From ff2086a1ff56e6c51bb7e8ad831478fec489cfea Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 10 Jan 2022 15:25:52 +0000 Subject: [PATCH] - refactored - use Armadillo for load/save of weight and training data git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@775 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Layer.cpp | 17 +++------- source/Layer.hpp | 55 +++++++++++++++++++++++++++++--- source/MainComponent.cpp | 8 ++--- source/MainComponent.hpp | 10 +++--- source/RbmComponent.cpp | 10 +++--- source/RbmComponent.hpp | 2 +- source/Stack.cpp | 68 ++++++++++++++++++++++++++-------------- source/Stack.hpp | 10 +++--- source/main.cpp | 10 +++--- 9 files changed, 125 insertions(+), 65 deletions(-) diff --git a/source/Layer.cpp b/source/Layer.cpp index c912788..1799bb1 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -26,7 +26,6 @@ Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibl , m_context(0, numContext) { cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl; - m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat"); } Layer::Layer(const Layer& orig) @@ -34,7 +33,6 @@ Layer::Layer(const Layer& orig) , next(nullptr) , prev(nullptr) , m_name(orig.m_name) -, m_weightsFile(orig.m_weightsFile) , m_id(orig.m_id) , m_numVisibleX(orig.m_numVisibleX) , m_numVisibleY(orig.m_numVisibleY) @@ -45,13 +43,10 @@ Layer::~Layer() { } + bool Layer::loadWeights(const string &prjname) { - string filename = m_weightsFile; - if (prjname.size() > 0) - { - filename = prjname + "." + m_weightsFile; - } + string filename = filePrefix(prjname) + ".weights.dat"; FILE *pFile = fopen(filename.c_str(),"r"); if (!pFile) @@ -108,11 +103,7 @@ bool Layer::saveWeights(const string &prjname) { int numHidden = m_bhv.n_elem; int numVisible = m_bv.n_elem; - string filename = m_weightsFile; - if (prjname.size() > 0) - { - filename = prjname + "." + m_weightsFile; - } + string filename = filePrefix(prjname) + ".weights.dat"; FILE *pFile = fopen(filename.c_str(),"w"); if (!pFile) @@ -153,7 +144,7 @@ Json::Value Layer::toJson() const Json::Value layer; layer["name"] = m_name; layer["id"] = (int)m_id; - layer["weights_file"] = m_weightsFile; + layer["weights_file"] = filePrefix("") + ".weights.dat"; layer["numVisibleX"] = (int)m_numVisibleX; layer["numVisibleY"] = (int)m_numVisibleY; layer["numHidden"] = (int)whv().n_cols; diff --git a/source/Layer.hpp b/source/Layer.hpp index 334303b..68e61cb 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -31,9 +31,7 @@ public: virtual ~Layer(); Json::Value toJson() const; - bool loadWeights(const std::string &prjname=""); - bool saveWeights(const std::string &prjname=""); - + void setBatch(arma::mat const &batch) { if (batch.n_rows > 0) @@ -83,6 +81,43 @@ public: return whv().n_cols; } + bool weightsLoad(std::string const &dir, std::string const &prj) + { + arma::mat w; + arma::mat bh; + arma::mat bv; + bool result = true; + result &= w.load(filePrefix(prj) + ".w.dat", arma::arma_ascii); + result &= bh.load(filePrefix(prj) + ".bh.dat", arma::arma_ascii); + result &= bv.load(filePrefix(prj) + ".bv.dat", arma::arma_ascii); + + if (result) + { + std::cout << "Layer " << m_id << ": Importing weights" << std::endl; + weightsAssign(w, bh, bv); + } + else + { + return loadWeights(prj); + } + return true; + } + + bool weightsSave(std::string const &dir, std::string const &prj) + { + bool result = true; + result &= whv().save(filePrefix(prj) + ".w.dat", arma::arma_ascii); + result &= bh().save(filePrefix(prj) + ".bh.dat", arma::arma_ascii); + result &= bv().save(filePrefix(prj) + ".bv.dat", arma::arma_ascii); + + if (result) + { + std::cout << "Layer " << m_id << ": Exporting weights" << std::endl; + } + + return result; + } + arma::mat trainingData(arma::mat const &batch) { arma::mat thisBatch = batch; @@ -131,12 +166,24 @@ public: private: std::string m_name; - std::string m_weightsFile; size_t m_id; size_t m_numVisibleX; size_t m_numVisibleY; size_t m_numContext; arma::mat m_context; + + // Compatibility + bool loadWeights(const std::string &prjname=""); + bool saveWeights(const std::string &prjname=""); + std::string filePrefix(const std::string &prjname) const + { + std::string filename = m_name + "." + std::to_string((int)m_id); + if (prjname.size() > 0) + { + filename = prjname + "." + filename; + } + return filename; + } }; diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp index e3534db..c1b00a1 100644 --- a/source/MainComponent.cpp +++ b/source/MainComponent.cpp @@ -814,7 +814,7 @@ void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged) m_pLayer = static_cast(m_stack->getLayer(index)); updateControls(); m_pLayer->redrawWeights(m_weightIndex); - if (m_stack->trainingData().n_rows > 0) + if (m_stack->trainingBatch().n_rows > 0) { m_pLayer->setTrainingData(trainingAt(m_trainingIndex)); } @@ -878,7 +878,7 @@ void MainComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails //[MiscUserCode] You can add your own definitions of your custom methods or any other code here... void MainComponent::save () { - m_stack->saveTraining(); + m_stack->saveTrainingBatch(); m_stack->saveWeights(); m_stack->save(); } @@ -899,7 +899,7 @@ const juce::String& MainComponent::getBaseDir() void MainComponent::run() { trainButton->setButtonText (TRANS("Stop")); - m_pLayer->setBatch(m_stack->trainingData()); + m_pLayer->setBatch(m_stack->trainingBatch()); m_pLayer->train(this); trainButton->setButtonText (TRANS("Train")); } @@ -916,7 +916,7 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status) } RbmComponent *pComp = static_cast(pRbm); - pComp->upPass(pComp->getTraining()); + pComp->upPass(pComp->getTrainingPattern()); pComp->redrawReconstruction(); pComp->redrawWeights(); diff --git a/source/MainComponent.hpp b/source/MainComponent.hpp index 4445dfc..3f2be79 100644 --- a/source/MainComponent.hpp +++ b/source/MainComponent.hpp @@ -98,12 +98,12 @@ private: void clearTraining() { patterSlider->setRange(0, 0, 1); - m_stack->trainingData().clear(); + m_stack->trainingBatch().clear(); } void loadTraining() { - m_stack->loadTraining(rbmNormalizeDataToggleButton->getToggleState()); + m_stack->loadTrainingBatch(rbmNormalizeDataToggleButton->getToggleState()); patterSlider->setRange(0, m_stack->numTraining()-1, 1); } @@ -128,16 +128,16 @@ private: { if (m_pLayer->context().is_empty()) { - m_pLayer->setBatch(m_stack->trainingData()); + m_pLayer->setBatch(m_stack->trainingBatch()); } if (!m_pLayer->context().is_empty()) { - return arma::join_rows(m_stack->trainingData().row(index), m_pLayer->context().row(index)); + return arma::join_rows(m_stack->trainingBatch().row(index), m_pLayer->context().row(index)); } else { - return m_stack->trainingData().row(index); + return m_stack->trainingBatch().row(index); } } diff --git a/source/RbmComponent.cpp b/source/RbmComponent.cpp index c434fba..fb8f77b 100644 --- a/source/RbmComponent.cpp +++ b/source/RbmComponent.cpp @@ -215,11 +215,11 @@ void RbmComponent::onDraw(DrawComponent &obj) } if (&obj == DrawVisibleTrain) { - upDownPass(getTraining()); + upDownPass(getTrainingPattern()); } if (&obj == DrawContextTrain) { - upDownPass(getTraining()); + upDownPass(getTrainingPattern()); } } @@ -249,14 +249,14 @@ void RbmComponent::buttonClicked(Button* buttonThatWasClicked) else if (buttonThatWasClicked == m_buttonCopyH2C) { DrawContextTrain->getData() = DrawHidden->getData(); - upDownPass(getTraining()); + upDownPass(getTrainingPattern()); } } void RbmComponent::redrawReconstruction() { RbmComponent *pComp = static_cast (root()); - pComp->upDownPass(pComp->getTraining()); + pComp->upDownPass(pComp->getTrainingPattern()); } void RbmComponent::gibbs(const arma::mat& vc) @@ -277,7 +277,7 @@ void RbmComponent::gibbs(const arma::mat& vc) DrawContextReconst->DrawData(); } -arma::mat RbmComponent::getTraining() const +arma::mat RbmComponent::getTrainingPattern() const { return arma::join_rows(DrawVisibleTrain->getData(), DrawContextTrain->getData()); } diff --git a/source/RbmComponent.hpp b/source/RbmComponent.hpp index 798117d..c966d38 100644 --- a/source/RbmComponent.hpp +++ b/source/RbmComponent.hpp @@ -73,7 +73,7 @@ public: ScopedPointer DrawVisibleTrain; ScopedPointer DrawHidden; ScopedPointer DrawContextTrain; - arma::mat getTraining() const; + arma::mat getTrainingPattern() const; arma::mat getReconst() const; void trainRedraw(const arma::mat& vc); void reconstRedraw(const arma::mat& vc); diff --git a/source/Stack.cpp b/source/Stack.cpp index 9ed3358..3aebeb7 100644 --- a/source/Stack.cpp +++ b/source/Stack.cpp @@ -175,7 +175,7 @@ bool Stack::loadWeights() Layer *pLayer = m_pLayers; while(pLayer) { - if (!pLayer->loadWeights(m_dir + "/" + m_name)) + if (!pLayer->weightsLoad(m_dir, m_name)) { return false; } @@ -189,7 +189,7 @@ bool Stack::saveWeights() Layer *pLayer = m_pLayers; while(pLayer) { - if (!pLayer->saveWeights(m_dir + "/" + m_name)) + if (!pLayer->weightsSave(m_dir, m_name)) { return false; } @@ -204,20 +204,20 @@ void Stack::train(Rbm::IListener* pListener) while(pLayer) { std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl; - pLayer->setBatch(m_trainingData); + pLayer->setBatch(m_trainingBatch); pLayer->train(pListener); pLayer = pLayer->next; } } -arma::mat& Stack::trainingData() +arma::mat& Stack::trainingBatch() { - return m_trainingData; + return m_trainingBatch; } -arma::mat Stack::trainingData(Layer* pThatLayer) +arma::mat Stack::trainingBatch(Layer* pThatLayer) { - arma::mat thisBatch = m_trainingData; + arma::mat thisBatch = m_trainingBatch; Layer *pLayer = m_pLayers; while (pLayer) { @@ -231,8 +231,19 @@ arma::mat Stack::trainingData(Layer* pThatLayer) return thisBatch; } -size_t Stack::loadTraining(bool doNormalize) +size_t Stack::loadTrainingBatch(bool doNormalize) { + { + std::string path = m_dir + "/" + m_name + ".training.mat"; + bool success = m_trainingBatch.load(path, arma::arma_ascii); + + if (success) + { + std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; + return m_trainingBatch.n_rows; + } + } + uint32_t numTraining = 0; uint32_t numVisible = 0; @@ -255,8 +266,8 @@ size_t Stack::loadTraining(bool doNormalize) { return 0; } - m_trainingData = arma::zeros(numTraining, numVisible); - + m_trainingBatch = arma::zeros(numTraining, numVisible); + uint32_t i, j; for (i=0; i < numTraining; i++) { @@ -266,7 +277,7 @@ size_t Stack::loadTraining(bool doNormalize) int result = fscanf(pFile, "%f", &v); if (result > 0) { - m_trainingData(i, j) = v; + m_trainingBatch(i, j) = v; } } } @@ -275,13 +286,24 @@ size_t Stack::loadTraining(bool doNormalize) if (doNormalize) { - m_trainingData = Rbm::normalize(m_trainingData); + m_trainingBatch = Rbm::normalize(m_trainingData); } return numTraining; } -size_t Stack::saveTraining() +size_t Stack::saveTrainingBatch() { + { + std::string path = m_dir + "/" + m_name + ".training.mat"; + bool success = m_trainingBatch.save(path, arma::arma_ascii); + + if (success) + { + std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n"; + return m_trainingBatch.n_rows; + } + } + std::string filename = m_dir + "/" + m_name + ".training.dat"; FILE *pFile = fopen(filename.c_str(), "w"); @@ -291,37 +313,37 @@ size_t Stack::saveTraining() return 0; } - fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_rows); - fprintf(pFile, "%u\n", (uint32_t)m_trainingData.n_cols); + fprintf(pFile, "%u\n", (uint32_t)m_trainingBatch.n_rows); + fprintf(pFile, "%u\n", (uint32_t)m_trainingBatch.n_cols); uint32_t i, j; - for (i=0; i < m_trainingData.n_rows; i++) + for (i=0; i < m_trainingBatch.n_rows; i++) { - for (j=0; j < m_trainingData.n_cols; j++) + for (j=0; j < m_trainingBatch.n_cols; j++) { - fprintf(pFile, "%3.6f\n", m_trainingData(i, j)); + fprintf(pFile, "%3.6f\n", m_trainingBatch(i, j)); } } fclose(pFile); - std::cout << "Saved " << m_trainingData.n_rows << " training samples\n"; - return m_trainingData.n_rows; + std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n"; + return m_trainingBatch.n_rows; } size_t Stack::numTraining() { - return m_trainingData.n_rows; + return m_trainingBatch.n_rows; } void Stack::addTraining(const arma::mat &toAdd) { - m_trainingData.insert_rows(m_trainingData.n_rows, toAdd); + m_trainingBatch.insert_rows(m_trainingBatch.n_rows, toAdd); } void Stack::delTraining(int index) { - m_trainingData.shed_row(index); + m_trainingBatch.shed_row(index); } diff --git a/source/Stack.hpp b/source/Stack.hpp index 36235a6..fa6442e 100644 --- a/source/Stack.hpp +++ b/source/Stack.hpp @@ -55,16 +55,16 @@ public: size_t numTraining(); void addTraining(const arma::mat &toAdd); void delTraining(int index); - size_t loadTraining(bool doNormalize=false); - size_t saveTraining(); - arma::mat& trainingData(); - arma::mat trainingData(Layer *pLayer); + size_t loadTrainingBatch(bool doNormalize=false); + size_t saveTrainingBatch(); + arma::mat& trainingBatch(); + arma::mat trainingBatch(Layer *pLayer); private: std::string m_dir; std::string m_name; Layer *m_pLayers; - arma::mat m_trainingData; + arma::mat m_trainingBatch; }; diff --git a/source/main.cpp b/source/main.cpp index 34b27e3..158af96 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -82,13 +82,13 @@ int main() RbmListener statusDisplay; Stack stack(project); - stack.loadTraining(); + stack.loadTrainingBatch(); - stack.addTraining(stack.trainingData().row(1)); - printf("There are %d training samples\n", (int)stack.trainingData().n_rows); + stack.addTraining(stack.trainingBatch().row(1)); + printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); stack.delTraining(0); - printf("There are %d training samples\n", (int)stack.trainingData().n_rows); + printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); #if 1 const int numLayers = 4; @@ -129,7 +129,7 @@ int main() stack.saveWeights(); Layer *layer = stack.getLayer(0); - arma::mat v = arma::randu(stack.trainingData().n_rows, layer->bv().n_elem); + arma::mat v = arma::randu(stack.trainingBatch().n_rows, layer->bv().n_elem); arma::mat h = layer->toHiddenProbs(v); arma::mat r = layer->toVisibleProbs(h); return 0;