- 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:
2022-01-11 16:46:01 +00:00
parent d18b45ef50
commit eacefc1e80
3 changed files with 55 additions and 9 deletions
+47 -2
View File
@@ -76,9 +76,54 @@ public:
return m_numVisibleY; 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); bool weightsLoad(std::string const &dir, std::string const &prj);
+2 -6
View File
@@ -349,13 +349,9 @@ void RbmComponent::upDownPass(const arma::mat& vc)
} }
else if (prev) 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) arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w)
+6 -1
View File
@@ -84,13 +84,13 @@ int main()
stack.loadTrainingBatch(); stack.loadTrainingBatch();
#if CREATE_TEST
stack.addTraining(stack.trainingBatch().row(1)); stack.addTraining(stack.trainingBatch().row(1));
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
stack.delTraining(0); stack.delTraining(0);
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
#if 1
const int numLayers = 4; const int numLayers = 4;
int i = 0; int i = 0;
Layer *lowerLayer = new Layer("Layer", i, 16, 16, 8); 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 v = arma::randu(stack.trainingBatch().n_rows, layer->bv().n_elem);
arma::mat h = layer->toHiddenProbs(v); arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h); arma::mat r = layer->toVisibleProbs(h);
v.print("v1");
r = layer->upDownPass(v);
r.print("v2");
return 0; return 0;
} }