cd_hinton_hid_binary: added raoBlackwell, gibbs-sampling
This commit is contained in:
+39
-7
@@ -175,20 +175,52 @@ void Rbm::cd_hinton_hid_binary(arma::mat const &v_data, arma::mat &dw, arma::mat
|
|||||||
{
|
{
|
||||||
// Start positive phase
|
// Start positive phase
|
||||||
arma::mat poshidprobs = prob(v_to_h(v_data));
|
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 poshidact = arma::sum(poshidprobs);
|
||||||
arma::mat posvisact = arma::sum(v_data);
|
arma::mat posvisact = arma::sum(v_data);
|
||||||
|
|
||||||
// End of positive phase
|
// End of positive phase
|
||||||
arma::mat poshidstates = sample(poshidprobs);
|
arma::mat poshidstates = sample(poshidprobs);
|
||||||
|
|
||||||
// Start negative phase
|
if (m_params.doRaoBlackwell)
|
||||||
arma::mat negdata = prob(h_to_v(poshidstates));
|
{
|
||||||
arma::mat neghidprobs = prob(v_to_h(negdata));
|
posprods = v_data.t() * poshidprobs;
|
||||||
arma::mat negprods = negdata.t() * neghidprobs;
|
}
|
||||||
arma::mat neghidact = arma::sum(neghidprobs);
|
else
|
||||||
arma::mat negvisact = arma::sum(negdata);
|
{
|
||||||
|
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
|
// Update weight deltas
|
||||||
dw = posprods - negprods;
|
dw = posprods - negprods;
|
||||||
dbv = posvisact - negvisact;
|
dbv = posvisact - negvisact;
|
||||||
|
|||||||
Reference in New Issue
Block a user