diff --git a/source/Rbm.cpp b/source/Rbm.cpp index ac33801..dd7f0d2 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -63,11 +63,10 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, size_t epoch; size_t gibbs; - size_t trainingSize = batch.n_rows; - status.trainingSizeRemain = trainingSize; + status.trainingSizeRemain = batch.n_rows; size_t batchRowIndex = 0; - double dProgress = 1.0/trainingSize; + double dProgress = 1.0/status.trainingSizeRemain; arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows)); arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols)); @@ -88,8 +87,8 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); status.trainingSizeRemain -= miniBatchSizeActual; batchRowIndex += miniBatchSizeActual; - double learning_rate = m_params.learningRate/std::min(miniBatchSizeActual, trainingSize); - double weight_decay = m_params.weightDecay/std::min(miniBatchSizeActual, trainingSize); + double learning_rate = m_params.learningRate/miniBatchSizeActual; + double weight_decay = m_params.weightDecay/miniBatchSizeActual; arma::mat vis_state(miniBatchSizeActual, m_w.n_rows); arma::mat vis_probs(miniBatchSizeActual, m_w.n_rows);