- added passes to layer
- removed redundancy git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@784 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+47
-2
@@ -76,9 +76,54 @@ public:
|
||||
return m_numVisibleY;
|
||||
}
|
||||
|
||||
int numHidden()
|
||||
arma::mat gibbsPass(arma::mat &vr)
|
||||
{
|
||||
return whv().n_cols;
|
||||
arma::mat h;
|
||||
for (int i=0; i < params().numGibbs; i++)
|
||||
{
|
||||
h = prob(v_to_h(vr));
|
||||
vr = prob(h_to_v(h));
|
||||
}
|
||||
return h;
|
||||
}
|
||||
|
||||
arma::mat upPass(arma::mat const &v)
|
||||
{
|
||||
arma::mat r = v;
|
||||
arma::mat h = gibbsPass(r);
|
||||
if (next)
|
||||
{
|
||||
next->upPass(h);
|
||||
}
|
||||
else
|
||||
{
|
||||
return h;
|
||||
}
|
||||
}
|
||||
|
||||
arma::mat downPass(arma::mat const &h)
|
||||
{
|
||||
arma::mat v = prob(h_to_v(h));
|
||||
if (prev)
|
||||
{
|
||||
return prev->downPass(v);
|
||||
}
|
||||
return v;
|
||||
}
|
||||
|
||||
arma::mat upDownPass(arma::mat const &v)
|
||||
{
|
||||
arma::mat r = v;
|
||||
arma::mat h = gibbsPass(r);
|
||||
if (next)
|
||||
{
|
||||
next->upDownPass(h);
|
||||
}
|
||||
else if (prev)
|
||||
{
|
||||
prev->downPass(r);
|
||||
}
|
||||
return r;
|
||||
}
|
||||
|
||||
bool weightsLoad(std::string const &dir, std::string const &prj);
|
||||
|
||||
@@ -349,13 +349,9 @@ void RbmComponent::upDownPass(const arma::mat& vc)
|
||||
}
|
||||
else if (prev)
|
||||
{
|
||||
if (prev)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
|
||||
pComp->downPass(getReconst());
|
||||
}
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
|
||||
pComp->downPass(getReconst());
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w)
|
||||
|
||||
+6
-1
@@ -84,13 +84,13 @@ int main()
|
||||
|
||||
stack.loadTrainingBatch();
|
||||
|
||||
#if CREATE_TEST
|
||||
stack.addTraining(stack.trainingBatch().row(1));
|
||||
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
|
||||
|
||||
stack.delTraining(0);
|
||||
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
|
||||
|
||||
#if 1
|
||||
const int numLayers = 4;
|
||||
int i = 0;
|
||||
Layer *lowerLayer = new Layer("Layer", i, 16, 16, 8);
|
||||
@@ -132,5 +132,10 @@ int main()
|
||||
arma::mat v = arma::randu(stack.trainingBatch().n_rows, layer->bv().n_elem);
|
||||
arma::mat h = layer->toHiddenProbs(v);
|
||||
arma::mat r = layer->toVisibleProbs(h);
|
||||
|
||||
v.print("v1");
|
||||
r = layer->upDownPass(v);
|
||||
r.print("v2");
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user