- crate and train Stack

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@586 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-26 09:49:12 +00:00
parent 9fd440e761
commit 91cc83d9ba
4 changed files with 63 additions and 9 deletions
+14 -5
View File
@@ -90,19 +90,28 @@ int main()
size_t numTraining = batch.n_rows;
size_t numVisibleX = 28;
size_t numVisibleY = 28;
size_t numHidden = 64;
size_t numHidden = 256;
printf("Loaded %d training samples\n", (int)numTraining);
for (int i=0; i < 8; i++)
int i=0;
RbmLayer *lowerLayer = new RbmLayer(project, i, numVisibleX, numVisibleY, numHidden, rbmParams);
stack.addLayer(lowerLayer);
numHidden >>= 1;
i++;
for (i; i < 4; i++)
{
RbmLayer *layer = new RbmLayer(project, i, numVisibleX, numVisibleY, numHidden, rbmParams);
RbmLayer *layer = new RbmLayer(project, i, lowerLayer->bh().n_elem, 1, numHidden, rbmParams);
lowerLayer = layer;
stack.addLayer(layer);
numHidden >>= 1;
}
stack.save(numTraining);
stack.train(batch, 1000, 100, &statusDisplay);
stack.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 = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h);