- 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
+1
View File
@@ -24,6 +24,7 @@ RbmLayer::RbmLayer(const string &prjname, size_t id, size_t numVisibleX, size_t
, m_numVisibleY(numVisibleY)
, m_rbm_params(params)
{
cout << "Create Layer " << m_prjname << "::" << to_string((int)m_id) << endl;
m_weightsFile = m_prjname + string(".weights.") + to_string((int)m_id) + string(".dat");
}
+1 -1
View File
@@ -50,7 +50,7 @@ void Stack::addLayer(RbmLayer *pOtherLayer)
}
}
const RbmLayer* Stack::getLayer(size_t id)
RbmLayer* Stack::getLayer(size_t id) const
{
RbmLayer *pLayer = m_pLayers;
while(pLayer)
+1 -1
View File
@@ -28,7 +28,7 @@ public:
virtual ~Stack();
void addLayer(RbmLayer *pLayer);
const RbmLayer* getLayer(size_t id);
RbmLayer* getLayer(size_t id) const;
void save(size_t numTraining);
private:
+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;
}