- RBM fixed progress indication
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@569 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+6
-8
@@ -47,7 +47,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
|||||||
size_t trainingSizeRemain = trainingSize;
|
size_t trainingSizeRemain = trainingSize;
|
||||||
size_t batchRowIndex = 0;
|
size_t batchRowIndex = 0;
|
||||||
|
|
||||||
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(miniBatchSize, trainingSize));
|
double dProgress = 1.0/trainingSize;
|
||||||
|
|
||||||
arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows));
|
arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows));
|
||||||
arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols));
|
arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols));
|
||||||
@@ -63,11 +63,10 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
|||||||
while (trainingSizeRemain)
|
while (trainingSizeRemain)
|
||||||
{
|
{
|
||||||
status.trainingSizeRemain = trainingSizeRemain;
|
status.trainingSizeRemain = trainingSizeRemain;
|
||||||
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
|
size_t miniBatchSizeActual = std::min(miniBatchSize, trainingSizeRemain);
|
||||||
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+toSlice-1);
|
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
||||||
trainingSizeRemain -= toSlice;
|
trainingSizeRemain -= miniBatchSizeActual;
|
||||||
batchRowIndex += toSlice;
|
batchRowIndex += miniBatchSizeActual;
|
||||||
size_t miniBatchSizeActual = miniBatch.n_rows;
|
|
||||||
double learning_rate = m_params.m_learningRate/std::min(miniBatchSizeActual, trainingSize);
|
double learning_rate = m_params.m_learningRate/std::min(miniBatchSizeActual, trainingSize);
|
||||||
double weight_decay = m_params.m_weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
double weight_decay = m_params.m_weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
||||||
|
|
||||||
@@ -150,8 +149,6 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
|||||||
m_bh += learning_rate*momentum_bias_h;
|
m_bh += learning_rate*momentum_bias_h;
|
||||||
m_w += learning_rate*momentum_weights;
|
m_w += learning_rate*momentum_weights;
|
||||||
|
|
||||||
status.progress += dProgress;
|
|
||||||
|
|
||||||
} // Number of epochs
|
} // Number of epochs
|
||||||
|
|
||||||
status.epoch = epoch;
|
status.epoch = epoch;
|
||||||
@@ -159,6 +156,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
|||||||
arma::mat diffErr_squared = diffErr % diffErr;
|
arma::mat diffErr_squared = diffErr % diffErr;
|
||||||
status.err = accu(diffErr_squared)/diffErr_squared.n_elem;
|
status.err = accu(diffErr_squared)/diffErr_squared.n_elem;
|
||||||
status.err_total = 0;
|
status.err_total = 0;
|
||||||
|
status.progress += dProgress*miniBatchSizeActual;
|
||||||
|
|
||||||
if (pListener)
|
if (pListener)
|
||||||
{
|
{
|
||||||
|
|||||||
+2
-2
@@ -25,12 +25,12 @@ public:
|
|||||||
{
|
{
|
||||||
Params()
|
Params()
|
||||||
: m_weightInit(0.01)
|
: m_weightInit(0.01)
|
||||||
, m_weightDecay(0.0)
|
, m_weightDecay(0.001)
|
||||||
, m_learningRate(0.1)
|
, m_learningRate(0.1)
|
||||||
, m_momentum(0.5)
|
, m_momentum(0.5)
|
||||||
, m_doRaoBlackwell(true)
|
, m_doRaoBlackwell(true)
|
||||||
, m_gibbsDoSampleVisible(false)
|
, m_gibbsDoSampleVisible(false)
|
||||||
, m_gibbsDoSampleHidden(false)
|
, m_gibbsDoSampleHidden(true)
|
||||||
, m_doSampleBatch(false)
|
, m_doSampleBatch(false)
|
||||||
, m_numGibbs(1)
|
, m_numGibbs(1)
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user