From f7acfc5bd094d5d3be6b0c95044fad7e362b9afe Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Thu, 24 Oct 2019 18:34:41 +0000 Subject: [PATCH] - RBM fixed progress indication git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@569 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Rbm.cpp | 14 ++++++-------- source/Rbm.hpp | 4 ++-- source/main.cpp | 2 +- 3 files changed, 9 insertions(+), 11 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 014fce5..d91a3a5 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -47,7 +47,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize, size_t trainingSizeRemain = trainingSize; 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_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) { status.trainingSizeRemain = trainingSizeRemain; - size_t toSlice = std::min(miniBatchSize, trainingSizeRemain); - arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+toSlice-1); - trainingSizeRemain -= toSlice; - batchRowIndex += toSlice; - size_t miniBatchSizeActual = miniBatch.n_rows; + size_t miniBatchSizeActual = std::min(miniBatchSize, trainingSizeRemain); + arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); + trainingSizeRemain -= miniBatchSizeActual; + batchRowIndex += miniBatchSizeActual; double learning_rate = m_params.m_learningRate/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_w += learning_rate*momentum_weights; - status.progress += dProgress; - } // Number of epochs 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; status.err = accu(diffErr_squared)/diffErr_squared.n_elem; status.err_total = 0; + status.progress += dProgress*miniBatchSizeActual; if (pListener) { diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 8fbe3f6..6bb61fb 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -25,12 +25,12 @@ public: { Params() : m_weightInit(0.01) - , m_weightDecay(0.0) + , m_weightDecay(0.001) , m_learningRate(0.1) , m_momentum(0.5) , m_doRaoBlackwell(true) , m_gibbsDoSampleVisible(false) - , m_gibbsDoSampleHidden(false) + , m_gibbsDoSampleHidden(true) , m_doSampleBatch(false) , m_numGibbs(1) { diff --git a/source/main.cpp b/source/main.cpp index d1ba959..d20a1b8 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -126,4 +126,4 @@ int main() arma::mat h = rbm.toHiddenProbs(v); arma::mat r = rbm.toVisibleProbs(h); return 0; -} \ No newline at end of file +}