diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 4f1db0c..0d404bc 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -139,7 +139,8 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) arma::mat momentum_bias_v(arma::zeros(1, m_bv.n_cols)); arma::mat momentum_bias_hv(arma::zeros(1, m_bhv.n_cols)); arma::mat penalty_weights = arma::zeros(m_whv.n_rows, m_whv.n_cols); - + arma::mat ctx = arma::zeros(1, numContext()); + int trainingSizeRemain = batch.n_rows; bool shouldAbort = false; @@ -158,8 +159,23 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) arma::mat v_probs(miniBatchSizeActual, m_bv.n_cols); arma::mat hid_states(miniBatchSizeActual, m_bhv.n_cols); #endif - arma::mat c_states(arma::zeros(miniBatchSizeActual, m_ctx.n_cols)); - arma::mat miniBatch(arma::join_rows(miniBatch_v, c_states)); + 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