From fe8142627d45a7d1ca86c8c2de9d4f0d81f2fa46 Mon Sep 17 00:00:00 2001 From: jens Date: Thu, 25 Jan 2024 22:45:50 +0100 Subject: [PATCH] cd_hinton_hid_binary: added raoBlackwell, gibbs-sampling --- source/Rbm.cpp | 48 ++++++++++++++++++++++++++++++++++++++++-------- 1 file changed, 40 insertions(+), 8 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 9800e3c..8e13ae2 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -175,20 +175,52 @@ void Rbm::cd_hinton_hid_binary(arma::mat const &v_data, arma::mat &dw, arma::mat { // Start positive phase arma::mat poshidprobs = prob(v_to_h(v_data)); - arma::mat posprods = v_data.t() * poshidprobs; + + arma::mat posprods; arma::mat poshidact = arma::sum(poshidprobs); arma::mat posvisact = arma::sum(v_data); - // End of positive phase + // End of positive phase arma::mat poshidstates = sample(poshidprobs); - // Start negative phase - arma::mat negdata = prob(h_to_v(poshidstates)); - arma::mat neghidprobs = prob(v_to_h(negdata)); - arma::mat negprods = negdata.t() * neghidprobs; - arma::mat neghidact = arma::sum(neghidprobs); - arma::mat negvisact = arma::sum(negdata); + if (m_params.doRaoBlackwell) + { + posprods = v_data.t() * poshidprobs; + } + else + { + posprods = v_data.t() * sample(poshidprobs); + } + arma::mat negprods; + arma::mat neghidact; + arma::mat negvisact; + for (int i=0; i < m_params.numGibbs; i++) + { + // Start negative phase + arma::mat negdata = prob(h_to_v(poshidstates)); + arma::mat neghidprobs; + if (m_params.gibbsDoSampleVisible) + { + neghidprobs = prob(v_to_h(sample(negdata))); + } + else + { + neghidprobs = prob(v_to_h(negdata)); + } + negprods = negdata.t() * neghidprobs; + neghidact = arma::sum(neghidprobs); + negvisact = arma::sum(negdata); + + if (m_params.gibbsDoSampleHidden) + { + poshidprobs = sample(neghidprobs); + } + else + { + poshidprobs = neghidprobs; + } + } // Update weight deltas dw = posprods - negprods; dbv = posvisact - negvisact;