- correct energy calc. for BB-RBMs
git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@39 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+7
-17
@@ -105,7 +105,6 @@ public:
|
||||
sumBiasH.fill(0);
|
||||
for (i=0; i < vt.getSize(); i++)
|
||||
{
|
||||
//t = (uint32_t)(0.5 + (vt.getSize()-1)*Noise_Uniform(&m_noise, 0.5));
|
||||
t = i;
|
||||
h.probsUpdateLogistic(vt[t], m_w, m_lambda, sigma);
|
||||
|
||||
@@ -221,24 +220,15 @@ public:
|
||||
return m_progress;
|
||||
}
|
||||
|
||||
double getEnergy(VisibleLayer &v, HiddenLayer &h)
|
||||
double getEnergy(const VectorXd& visible, const VectorXd& hidden)
|
||||
{
|
||||
uint32_t i, j;
|
||||
double energy;
|
||||
|
||||
energy = -v.getEnergy(m_w) - h.getEnergy(m_w);
|
||||
energy = m_w.visibleBias().transpose() * visible;
|
||||
energy += m_w.hiddenBias().transpose() * hidden;
|
||||
energy += visible.transpose() * m_w.weights() * hidden;
|
||||
|
||||
for (i=0; i < h.getNumUnits(); i++)
|
||||
{
|
||||
for (j=0; j < v.getNumUnits(); j++)
|
||||
{
|
||||
// energy -= v.getStates()[j] * h.getStates()[i] * m_w.getWeights()[i][j];
|
||||
}
|
||||
}
|
||||
|
||||
// ToDo: make this correct
|
||||
// energy -= (v.states().transpose() * h.states()); // * m_w.weights();
|
||||
return energy;
|
||||
return -energy/(m_sigma*m_sigma);
|
||||
}
|
||||
|
||||
void prob(LayerArray<VisibleLayer> &vts)
|
||||
@@ -278,11 +268,11 @@ public:
|
||||
z = 0;
|
||||
for (j=0; j < vts.getSize(); j++)
|
||||
{
|
||||
z += exp(-getEnergy(vts.getAt(j), h[i]));
|
||||
z += exp(-getEnergy(vts.getAt(j).states(), h[i].states()));
|
||||
}
|
||||
for (j=0; j < vts.getSize(); j++)
|
||||
{
|
||||
p = exp(-getEnergy(vts.getAt(j), h[i]))/z;
|
||||
p = exp(-getEnergy(vts.getAt(j).states(), h[i].states()))/z;
|
||||
cout << p << endl;
|
||||
}
|
||||
cout << endl;
|
||||
|
||||
Reference in New Issue
Block a user