cd_hinton_hid_binary: added raoBlackwell, gibbs-sampling

This commit is contained in:
2024-01-25 22:45:50 +01:00
parent 34b20f4fb3
commit fe8142627d
+40 -8
View File
@@ -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;