- add more log info

- create layers dynamically in main

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@583 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-26 06:16:18 +00:00
parent d16f0de8fa
commit f50607ef74
4 changed files with 14 additions and 16 deletions
+11 -14
View File
@@ -83,6 +83,7 @@ int main()
RbmListener statusDisplay;
Rbm::Params rbmParams;
Stack stack(project);
arma::mat batch = loadTraining(project + string(".training.dat"));
@@ -93,21 +94,17 @@ int main()
printf("Loaded %d training samples\n", (int)numTraining);
RbmLayer layer0(project, 0, numVisibleX, numVisibleY, numHidden, rbmParams);
RbmLayer layer1(project, 1, numVisibleX, numVisibleY, numHidden, rbmParams);
RbmLayer layer2(project, 2, numVisibleX, numVisibleY, numHidden, rbmParams);
RbmLayer layer3(project, 3, numVisibleX, numVisibleY, numHidden, rbmParams);
Stack stack(project);
stack.addLayer(&layer0);
stack.addLayer(&layer1);
stack.addLayer(&layer2);
stack.addLayer(&layer3);
for (int i=0; i < 8; i++)
{
RbmLayer *layer = new RbmLayer(project, i, numVisibleX, numVisibleY, numHidden, rbmParams);
stack.addLayer(layer);
}
stack.save(numTraining);
layer0.train(batch, 1000, 100, &statusDisplay);
layer0.saveWeights();
RbmLayer *layer = stack.getLayer(0);
layer->train(batch, 1000, 100, &statusDisplay);
layer->saveWeights();
arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY);
arma::mat h = layer0.toHiddenProbs(v);
arma::mat r = layer0.toVisibleProbs(h);
arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h);
return 0;
}