diff --git a/source/Layer.hpp b/source/Layer.hpp index 9e6243c..a5998dc 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -37,8 +37,24 @@ public: void train(arma::mat const &batch, IListener *pListener=nullptr) { if (batch.n_rows > 0) - { - Rbm::train(trainingData(batch), pListener); + { + if (numContext()) + { + arma::mat c_states = arma::zeros(batch.n_rows, numContext()); + arma::mat ctx = arma::zeros(1, numContext()); + for (int i=1; i < batch.n_rows; i++) + { + arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx))); + ctx = h; + c_states.row(i) = ctx; + } + arma::mat batch_with_ctx = arma::join_rows(batch, c_states); + Rbm::train(trainingData(batch_with_ctx), pListener); + } + else + { + Rbm::train(trainingData(batch), pListener); + } } } diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 87a8ee8..948ea15 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -147,30 +147,13 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) while (trainingSizeRemain && !shouldAbort) { int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain); - arma::mat miniBatch_v = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); + arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); trainingSizeRemain -= miniBatchSizeActual; batchRowIndex += miniBatchSizeActual; int scaler = std::min(m_params.miniBatchSize, (int)batch.n_rows); double learning_rate = m_params.learningRate/scaler; double weight_decay = m_params.weightDecay/scaler; - arma::mat miniBatch; - if (numContext() > 0) - { - arma::mat c_states(arma::zeros(miniBatchSizeActual, numContext())); - - for (int i=1; i < miniBatchSizeActual; i++) - { - arma::mat h = prob(v_to_h(arma::join_rows(miniBatch_v.row(i-1), ctx))); - ctx = h; - c_states.row(i) = ctx; - } - miniBatch = (arma::join_rows(miniBatch_v, c_states)); - } - else - { - miniBatch = miniBatch_v; - } arma::mat v_states(miniBatch); // Create hidden layer base on training data @@ -221,7 +204,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) } // number of mini batches -#if FIXED_BATCH arma::mat diffErr = batch - prob(h_to_v(prob(v_to_h(batch)))); arma::mat diffErr_squared = diffErr % diffErr; status.err_total = accu(diffErr_squared)/diffErr_squared.n_elem; @@ -230,7 +212,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) { pListener->onProgress(this, status); } -#endif } arma::mat Rbm::prob(const arma::mat &src) diff --git a/source/Rbm.hpp b/source/Rbm.hpp index c7cc168..c5389b7 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -144,7 +144,7 @@ public: arma::mat toHiddenProbs(const arma::mat &visible) const { - return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx))); + return Rbm::prob(v_to_h(visible)); } arma::mat toVisibleProbs(const arma::mat &hidden) const