- 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:
@@ -24,6 +24,7 @@ RbmLayer::RbmLayer(const string &prjname, size_t id, size_t numVisibleX, size_t
|
|||||||
, m_numVisibleY(numVisibleY)
|
, m_numVisibleY(numVisibleY)
|
||||||
, m_rbm_params(params)
|
, 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");
|
m_weightsFile = m_prjname + string(".weights.") + to_string((int)m_id) + string(".dat");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -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;
|
RbmLayer *pLayer = m_pLayers;
|
||||||
while(pLayer)
|
while(pLayer)
|
||||||
|
|||||||
+1
-1
@@ -28,7 +28,7 @@ public:
|
|||||||
virtual ~Stack();
|
virtual ~Stack();
|
||||||
|
|
||||||
void addLayer(RbmLayer *pLayer);
|
void addLayer(RbmLayer *pLayer);
|
||||||
const RbmLayer* getLayer(size_t id);
|
RbmLayer* getLayer(size_t id) const;
|
||||||
void save(size_t numTraining);
|
void save(size_t numTraining);
|
||||||
|
|
||||||
private:
|
private:
|
||||||
|
|||||||
+11
-14
@@ -83,6 +83,7 @@ int main()
|
|||||||
|
|
||||||
RbmListener statusDisplay;
|
RbmListener statusDisplay;
|
||||||
Rbm::Params rbmParams;
|
Rbm::Params rbmParams;
|
||||||
|
Stack stack(project);
|
||||||
|
|
||||||
arma::mat batch = loadTraining(project + string(".training.dat"));
|
arma::mat batch = loadTraining(project + string(".training.dat"));
|
||||||
|
|
||||||
@@ -93,21 +94,17 @@ int main()
|
|||||||
|
|
||||||
printf("Loaded %d training samples\n", (int)numTraining);
|
printf("Loaded %d training samples\n", (int)numTraining);
|
||||||
|
|
||||||
RbmLayer layer0(project, 0, numVisibleX, numVisibleY, numHidden, rbmParams);
|
for (int i=0; i < 8; i++)
|
||||||
RbmLayer layer1(project, 1, numVisibleX, numVisibleY, numHidden, rbmParams);
|
{
|
||||||
RbmLayer layer2(project, 2, numVisibleX, numVisibleY, numHidden, rbmParams);
|
RbmLayer *layer = new RbmLayer(project, i, numVisibleX, numVisibleY, numHidden, rbmParams);
|
||||||
RbmLayer layer3(project, 3, numVisibleX, numVisibleY, numHidden, rbmParams);
|
stack.addLayer(layer);
|
||||||
Stack stack(project);
|
}
|
||||||
stack.addLayer(&layer0);
|
|
||||||
stack.addLayer(&layer1);
|
|
||||||
stack.addLayer(&layer2);
|
|
||||||
stack.addLayer(&layer3);
|
|
||||||
stack.save(numTraining);
|
stack.save(numTraining);
|
||||||
|
RbmLayer *layer = stack.getLayer(0);
|
||||||
layer0.train(batch, 1000, 100, &statusDisplay);
|
layer->train(batch, 1000, 100, &statusDisplay);
|
||||||
layer0.saveWeights();
|
layer->saveWeights();
|
||||||
arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY);
|
arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY);
|
||||||
arma::mat h = layer0.toHiddenProbs(v);
|
arma::mat h = layer->toHiddenProbs(v);
|
||||||
arma::mat r = layer0.toVisibleProbs(h);
|
arma::mat r = layer->toVisibleProbs(h);
|
||||||
return 0;
|
return 0;
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user