- 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;
|
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);
|
||||||
|
|||||||
@@ -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
@@ -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;
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user