diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 775a4c2..b4d7332 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -59,11 +59,12 @@ Json::Value Rbm::toJson() const void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener) { + Status status; size_t epoch; size_t gibbs; size_t trainingSize = batch.n_rows; - size_t trainingSizeRemain = trainingSize; + status.trainingSizeRemain = trainingSize; size_t batchRowIndex = 0; double dProgress = 1.0/trainingSize; @@ -76,15 +77,16 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, arma::mat momentum_bias_h(arma::zeros(1, m_w.n_cols)); arma::mat penalty_weights = arma::zeros(m_w.n_rows, m_w.n_cols); - Status status; - - status.progress = 0; - while (trainingSizeRemain) + if (pListener) { - status.trainingSizeRemain = trainingSizeRemain; - size_t miniBatchSizeActual = std::min(miniBatchSize, trainingSizeRemain); + pListener->onProgress(status); + } + + while (status.trainingSizeRemain) + { + size_t miniBatchSizeActual = std::min(miniBatchSize, status.trainingSizeRemain); arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); - trainingSizeRemain -= miniBatchSizeActual; + 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); diff --git a/source/Rbm.hpp b/source/Rbm.hpp index d72df94..0671f43 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -81,6 +81,16 @@ public: struct Status { + Status() + : epoch(0) + , trainingSizeRemain(0) + , progress(0) + , err(-1.0) + , err_total(1-0) + , L1(-1.0) + , L2(-1.0) + { + } size_t epoch; size_t trainingSizeRemain; double progress;