cd_hinton_hid_binary: added raoBlackwell, gibbs-sampling
This commit is contained in:
+40
-8
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user