From 711ad3057ebe0aee995e30d33891048a46657e51 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 28 Oct 2019 19:08:30 +0000 Subject: [PATCH] - Rbm: cleaned up, learning_rate and weight_decay normalized to miniBatchSizeActual git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@590 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Rbm.cpp | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) 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);