- 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
+2 -1
View File
@@ -12,8 +12,9 @@ CXXFLAGS += -std=c++11
CXXFLAGS_debug := ${CXXFLAGS} -O0 -g CXXFLAGS_debug := ${CXXFLAGS} -O0 -g
CXXFLAGS_release := ${CXXFLAGS} -O2 CXXFLAGS_release := ${CXXFLAGS} -O2
DEFINES := -DARMA_OPENMP_THREADS=1
${BUILD_DIR}/main.elf: ${BUILD_DIR} ${SRCS} ${BUILD_DIR}/main.elf: ${BUILD_DIR} ${SRCS}
g++ ${CXXFLAGS_${CONFIG}} ${SRCS} -o $@ ${LIBS} g++ ${CXXFLAGS_${CONFIG}} ${DEFINES} ${SRCS} -o $@ ${LIBS}
${BUILD_DIR}: ${BUILD_DIR}:
mkdir -p $@ mkdir -p $@
+41 -2
View File
@@ -50,12 +50,12 @@ void Stack::addLayer(RbmLayer *pOtherLayer)
} }
} }
RbmLayer* Stack::getLayer(size_t id) const RbmLayer* Stack::getLayer(size_t layerId) const
{ {
RbmLayer *pLayer = m_pLayers; RbmLayer *pLayer = m_pLayers;
while(pLayer) while(pLayer)
{ {
if (pLayer->id() == id) if (pLayer->id() == layerId)
{ {
return pLayer; return pLayer;
} }
@@ -88,4 +88,43 @@ void Stack::save(size_t numTraining)
ofs << writer.write(project); ofs << writer.write(project);
} }
void Stack::saveWeights()
{
RbmLayer *pLayer = m_pLayers;
while(pLayer)
{
pLayer->saveWeights();
}
}
void Stack::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener)
{
RbmLayer *pLayer = m_pLayers;
while(pLayer)
{
train(pLayer->id(), batch, miniBatchSize, numEpochs, pListener);
pLayer = pLayer->upper;
}
}
void Stack::train(size_t layerId, const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener)
{
arma::mat thisBatch = batch;
RbmLayer *pLayer = m_pLayers;
while(pLayer)
{
if (pLayer->id() == layerId)
{
break;
}
thisBatch = pLayer->toHiddenProbs(thisBatch);
pLayer = pLayer->upper;
}
if (pLayer)
{
std::cout << m_prjname << ": " << " Training of layer " << std::to_string(layerId) << std::endl;
pLayer->train(thisBatch, miniBatchSize, numEpochs, pListener);
}
}
+6 -1
View File
@@ -28,8 +28,12 @@ public:
virtual ~Stack(); virtual ~Stack();
void addLayer(RbmLayer *pLayer); void addLayer(RbmLayer *pLayer);
RbmLayer* getLayer(size_t id) const; RbmLayer* getLayer(size_t layerId) const;
void train(size_t layerId, const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener);
void train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener);
void save(size_t numTraining); void save(size_t numTraining);
void saveWeights();
private: private:
const std::string &m_prjname; const std::string &m_prjname;
@@ -37,5 +41,6 @@ private:
}; };
#endif /* STACK_HPP */ #endif /* STACK_HPP */
+14 -5
View File
@@ -90,19 +90,28 @@ int main()
size_t numTraining = batch.n_rows; size_t numTraining = batch.n_rows;
size_t numVisibleX = 28; size_t numVisibleX = 28;
size_t numVisibleY = 28; size_t numVisibleY = 28;
size_t numHidden = 64; size_t numHidden = 256;
printf("Loaded %d training samples\n", (int)numTraining); 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); stack.addLayer(layer);
numHidden >>= 1;
} }
stack.save(numTraining); stack.save(numTraining);
stack.train(batch, 1000, 100, &statusDisplay);
stack.saveWeights();
RbmLayer *layer = stack.getLayer(0); RbmLayer *layer = stack.getLayer(0);
layer->train(batch, 1000, 100, &statusDisplay);
layer->saveWeights();
arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY); arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY);
arma::mat h = layer->toHiddenProbs(v); arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h); arma::mat r = layer->toVisibleProbs(h);