From 56335914ea1bf7c8637addd9c03df7ef2cfecf32 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sun, 9 Jan 2022 08:34:59 +0000 Subject: [PATCH] - train with context git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@762 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Rbm.cpp | 22 +++++++++++++++++++--- 1 file changed, 19 insertions(+), 3 deletions(-) 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