From 9b2861e3d83b788f316fc5acfa9ab73dfc0e7c58 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 17 Jan 2022 20:34:44 +0000 Subject: [PATCH] - refactored git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@825 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/AStack.cpp | 65 +++++++++++++++++++++++++++++++++++++++++++ source/AStack.hpp | 10 ++++++- source/DeepStack.cpp | 66 -------------------------------------------- source/DeepStack.hpp | 9 ------ source/poet.cpp | 15 +++++----- 5 files changed, 81 insertions(+), 84 deletions(-) diff --git a/source/AStack.cpp b/source/AStack.cpp index 21407d1..bfdbed4 100644 --- a/source/AStack.cpp +++ b/source/AStack.cpp @@ -157,3 +157,68 @@ arma::mat AStack::trainingBatchFrom(size_t layerId, const arma::mat& batch) return thisBatch; } +arma::mat& AStack::trainingBatch() +{ + return m_trainingBatch; +} + +size_t AStack::loadTrainingBatch(const std::string &dir, bool doNormalize) +{ + std::string filename = dir + "/" + m_name + ".training.dat"; + std::string path = dir + "/" + m_name + ".training.dat"; + bool success = m_trainingBatch.load(filename, arma::arma_ascii); + + if (success) + { + if (doNormalize) + { + m_trainingBatch = Rbm::normalize(m_trainingBatch); + } + std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; + } + + // Migrate context part to training data + size_t numTraining = m_trainingBatch.n_rows; + size_t numVisible = getLayer(0)->numVisible(); + + if (m_trainingBatch.n_cols < numVisible) + { + size_t diff = numVisible - m_trainingBatch.n_cols; + arma::mat training_with_ctx = arma::join_rows(m_trainingBatch, arma::zeros(numTraining, diff)); + m_trainingBatch = training_with_ctx; + } + else if (m_trainingBatch.n_cols > numVisible) + { + m_trainingBatch = m_trainingBatch.submat(0, 0, numTraining-1, numVisible-1); + } + + return m_trainingBatch.n_rows; +} + +size_t AStack::saveTrainingBatch(const std::string &dir) +{ + std::string filename = dir + "/" + m_name + ".training.dat"; + bool success = m_trainingBatch.save(filename, arma::arma_ascii); + + if (success) + { + std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n"; + } + + return m_trainingBatch.n_rows; +} + +size_t AStack::numTraining() +{ + return m_trainingBatch.n_rows; +} + +void AStack::addTraining(const arma::mat &toAdd) +{ + m_trainingBatch.insert_rows(m_trainingBatch.n_rows, toAdd); +} + +void AStack::delTraining(int index) +{ + m_trainingBatch.shed_row(index); +} diff --git a/source/AStack.hpp b/source/AStack.hpp index 0b95715..54c190a 100644 --- a/source/AStack.hpp +++ b/source/AStack.hpp @@ -68,13 +68,21 @@ public: virtual void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) = 0; virtual void train(const arma::mat& batch, Rbm::IListener* pListener) = 0; arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch); - + + size_t numTraining(); + void addTraining(const arma::mat &toAdd); + void delTraining(int index); + size_t loadTrainingBatch(const std::string &dir, bool doNormalize=false); + size_t saveTrainingBatch(const std::string &dir); + arma::mat& trainingBatch(); + protected: StackType m_type; std::string m_name; Layer *m_pLayers; private: + arma::mat m_trainingBatch; }; diff --git a/source/DeepStack.cpp b/source/DeepStack.cpp index 8b54d2e..e8337db 100644 --- a/source/DeepStack.cpp +++ b/source/DeepStack.cpp @@ -52,72 +52,6 @@ void DeepStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pL pLayer->train(thisBatch, pListener); } -arma::mat& DeepStack::trainingBatch() -{ - return m_trainingBatch; -} - -size_t DeepStack::loadTrainingBatch(const std::string &dir, bool doNormalize) -{ - std::string filename = dir + "/" + m_name + ".training.dat"; - std::string path = dir + "/" + m_name + ".training.dat"; - bool success = m_trainingBatch.load(filename, arma::arma_ascii); - - if (success) - { - if (doNormalize) - { - m_trainingBatch = Rbm::normalize(m_trainingBatch); - } - std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; - } - - // Migrate context part to training data - size_t numTraining = m_trainingBatch.n_rows; - size_t numVisible = getLayer(0)->numVisible(); - - if (m_trainingBatch.n_cols < numVisible) - { - size_t diff = numVisible - m_trainingBatch.n_cols; - arma::mat training_with_ctx = arma::join_rows(m_trainingBatch, arma::zeros(numTraining, diff)); - m_trainingBatch = training_with_ctx; - } - else if (m_trainingBatch.n_cols > numVisible) - { - m_trainingBatch = m_trainingBatch.submat(0, 0, numTraining-1, numVisible-1); - } - - return m_trainingBatch.n_rows; -} - -size_t DeepStack::saveTrainingBatch(const std::string &dir) -{ - std::string filename = dir + "/" + m_name + ".training.dat"; - bool success = m_trainingBatch.save(filename, arma::arma_ascii); - - if (success) - { - std::cout << "Saved " << m_trainingBatch.n_rows << " training samples\n"; - } - - return m_trainingBatch.n_rows; -} - -size_t DeepStack::numTraining() -{ - return m_trainingBatch.n_rows; -} - -void DeepStack::addTraining(const arma::mat &toAdd) -{ - m_trainingBatch.insert_rows(m_trainingBatch.n_rows, toAdd); -} - -void DeepStack::delTraining(int index) -{ - m_trainingBatch.shed_row(index); -} - arma::mat DeepStack::upPass(size_t layerId, const arma::mat& v) { arma::mat h = arma::zeros(0,0); diff --git a/source/DeepStack.hpp b/source/DeepStack.hpp index 99e3373..a25be29 100644 --- a/source/DeepStack.hpp +++ b/source/DeepStack.hpp @@ -35,15 +35,6 @@ public: arma::mat downPass(size_t layerId, arma::mat const &h); arma::mat upDownPass(size_t layerId, arma::mat const &v); - size_t numTraining(); - void addTraining(const arma::mat &toAdd); - void delTraining(int index); - size_t loadTrainingBatch(const std::string &dir, bool doNormalize=false); - size_t saveTrainingBatch(const std::string &dir); - arma::mat& trainingBatch(); - -private: - arma::mat m_trainingBatch; }; diff --git a/source/poet.cpp b/source/poet.cpp index 93ff39c..41eb5fd 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -10,6 +10,7 @@ #include "Rbm.hpp" #include "Layer.hpp" #include "DeepStack.hpp" +#include "StackCreator.hpp" using namespace std; using namespace arma; @@ -199,21 +200,19 @@ int main() return 0; #endif - DeepStack stack(".", "poet5"); - // Load project - stack.load(); + AStack *stack = StackCreator::fromFile(".", "poet5"); // Load weights - stack.loadWeights(); + stack->loadWeights("."); // Load training - stack.loadTrainingBatch(); + stack->loadTrainingBatch("."); - Layer *layer = stack.getLayer(0); - int numTraining = stack.trainingBatch().n_rows; + Layer *layer = stack->getLayer(0); + int numTraining = stack->trainingBatch().n_rows; - arma::mat t = stack.trainingBatch(); + arma::mat t = stack->trainingBatch(); arma::mat h = layer->toHiddenProbs(t); arma::mat r = layer->toVisibleProbs(h);