- Rbm: fixed parameter import

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@607 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-06 06:24:48 +00:00
parent 3ef7b07f5e
commit b20604e4f4
4 changed files with 21 additions and 9 deletions
+14 -4
View File
@@ -89,20 +89,29 @@ int main()
size_t numTraining = batch.n_rows;
printf("Loaded %d training samples\n", (int)numTraining);
#if 1
#if 0
int i = 0;
Layer *lowerLayer = new Layer("Layer", i, 16, 16, 8);
stack.addLayer(lowerLayer);
for (++i; i < 1; i++)
for (++i; i < 4; i++)
{
Layer *layer = new Layer("Layer", i, lowerLayer->bh().n_elem, 1, lowerLayer->bh().n_elem >> 1);
lowerLayer = layer;
stack.addLayer(layer);
}
Layer *layer = stack.getLayer(0);
layer->params().learningRate = 0.02;
stack.getLayer(0)->params().learningRate = 0.04;
stack.getLayer(0)->params().numEpochs = 1000;
stack.getLayer(1)->params().learningRate = 0.03;
stack.getLayer(1)->params().numEpochs = 500;
stack.getLayer(2)->params().learningRate = 0.02;
stack.getLayer(2)->params().numEpochs = 250;
stack.getLayer(3)->params().learningRate = 0.01;
stack.getLayer(3)->params().numEpochs = 125;
// Save project
stack.save();
@@ -127,6 +136,7 @@ int main()
// Save weights
stack.saveWeights();
Layer *layer = stack.getLayer(0);
arma::mat v = arma::randu(numTraining, layer->bv().n_elem);
arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h);