diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 10f906a..2523c2e 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -132,12 +132,15 @@ public: private: const Params &m_params; - arma::mat m_w; - arma::mat m_bh; - arma::mat m_bv; arma::mat sample(arma::mat const &src); static arma::mat probsLogistic(arma::mat const &src); void uniform(arma::mat &srcDst, double mu=0.0, double stdDev=1.0); + +protected: + arma::mat m_w; + arma::mat m_bh; + arma::mat m_bv; + }; #endif /* RBM_HPP */ diff --git a/source/RbmLayer.cpp b/source/RbmLayer.cpp index 2f63e3b..1a2c299 100644 --- a/source/RbmLayer.cpp +++ b/source/RbmLayer.cpp @@ -45,41 +45,99 @@ RbmLayer::~RbmLayer() { } -void RbmLayer::saveWeights() +bool RbmLayer::loadWeights() { - FILE *pFile = fopen(m_weightsFile.c_str(), "w"); + int numVisibleX; + int numVisibleY; + int numHidden; + FILE *pFile; + + pFile = fopen(m_weightsFile.c_str(),"r"); if (!pFile) { - std::cout << "Could not open " << m_weightsFile << "!" << std::endl; - return; + std::cout << "loadWeights(): Could not open " << m_weightsFile << " for reading!" << std::endl; + return false; } - size_t numHidden = bh().n_elem; - size_t numVisible = bh().n_elem; - fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, (int)numHidden); - - uint32_t i, j; + int result = fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden); + if (result < 0) + { + return false; + } + size_t numVisible = numVisibleX*numVisibleY; + + int i, j; + float v; for (i=0; i < numVisible; i++) { - fprintf(pFile, "%3.6f\n", bv()(i)); + result = fscanf(pFile, "%f", &v); + if (result > 0) + { + m_bv(i) = v; + } } for (i=0; i < numHidden; i++) { - fprintf(pFile, "%3.6f\n", bh()(i)); + result = fscanf(pFile, "%f", &v); + if (result > 0) + { + m_bh(i) = v; + } } for (i=0; i < numVisible; i++) { for (j=0; j < numHidden; j++) { - fprintf(pFile, "%3.6f ", w()(i,j)); + + result = fscanf(pFile, "%f", &v); + if (result > 0) + { + m_w(i, j) = v; + } + } + } + fclose(pFile); + + return true; +} + +bool RbmLayer::saveWeights() +{ + FILE *pFile = fopen(m_weightsFile.c_str(), "w"); + + if (!pFile) + { + std::cout << "saveWeights(): Could not open " << m_weightsFile << " for writing!" << std::endl; + return false; + } + + int numHidden = m_bh.n_elem; + int numVisible = m_bv.n_elem; + fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, numHidden); + + int i, j; + + for (i=0; i < numVisible; i++) + { + fprintf(pFile, "%3.6f\n", m_bv(i)); + } + for (i=0; i < numHidden; i++) + { + fprintf(pFile, "%3.6f\n", m_bh(i)); + } + for (i=0; i < numVisible; i++) + { + for (j=0; j < numHidden; j++) + { + fprintf(pFile, "%3.6f ", m_w(i,j)); } fprintf(pFile, "\n"); } - fclose(pFile); + return true; } Json::Value RbmLayer::toJson() const diff --git a/source/RbmLayer.hpp b/source/RbmLayer.hpp index ef45256..80a9dbe 100644 --- a/source/RbmLayer.hpp +++ b/source/RbmLayer.hpp @@ -32,7 +32,8 @@ public: virtual ~RbmLayer(); Json::Value toJson() const; - void saveWeights(); + bool loadWeights(); + bool saveWeights(); arma::mat up_pass(const arma::mat& hidden); arma::mat down_pass(const arma::mat& visible); diff --git a/source/Stack.cpp b/source/Stack.cpp index dd73f6e..129a859 100644 --- a/source/Stack.cpp +++ b/source/Stack.cpp @@ -65,7 +65,7 @@ RbmLayer* Stack::getLayer(size_t layerId) const } -void Stack::save(size_t numTraining) +bool Stack::save(size_t numTraining) { std::cout << "Exporting Project " << m_prjname << std::endl; ofstream ofs(m_prjname + string(".prj")); @@ -86,16 +86,36 @@ void Stack::save(size_t numTraining) project["stack"]["layers"] = layers; ofs << writer.write(project); + + return true; } -void Stack::saveWeights() +bool Stack::loadWeights() { RbmLayer *pLayer = m_pLayers; while(pLayer) { - pLayer->saveWeights(); + if (!pLayer->loadWeights()) + { + return false; + } pLayer = pLayer->upper; } + return true; +} + +bool Stack::saveWeights() +{ + RbmLayer *pLayer = m_pLayers; + while(pLayer) + { + if (!pLayer->saveWeights()) + { + return false; + } + pLayer = pLayer->upper; + } + return true; } void Stack::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener) diff --git a/source/Stack.hpp b/source/Stack.hpp index 458b8f5..1e579dd 100644 --- a/source/Stack.hpp +++ b/source/Stack.hpp @@ -32,8 +32,9 @@ public: 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 saveWeights(); + bool save(size_t numTraining); + bool loadWeights(); + bool saveWeights(); private: const std::string &m_prjname; diff --git a/source/main.cpp b/source/main.cpp index 7d85b4b..c853e6e 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -107,8 +107,17 @@ int main() stack.addLayer(layer); numHidden >>= 1; } + + // Save project stack.save(numTraining); + + // Load weights + stack.loadWeights(); + + // Train stack stack.train(batch, 1000, 100, &statusDisplay); + + // Save weights stack.saveWeights(); RbmLayer *layer = stack.getLayer(0);