From 17770f402ce627ea145f951979884c437a03315f Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Tue, 18 Jan 2022 08:36:10 +0000 Subject: [PATCH] - integrated RnnStack git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@826 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/AStack.hpp | 6 +++- source/DeepStack.cpp | 16 +++++++++ source/DeepStack.hpp | 1 + source/Layer.cpp | 17 --------- source/Layer.hpp | 1 - source/MainComponent.cpp | 3 +- source/RnnStack.cpp | 76 ++++++++++++++++++++++++++++++++++++++-- source/RnnStack.hpp | 14 ++++++-- source/StackCreator.cpp | 6 ++-- source/poet.cpp | 24 ++++++++++--- 10 files changed, 130 insertions(+), 34 deletions(-) diff --git a/source/AStack.hpp b/source/AStack.hpp index 54c190a..d1d06e6 100644 --- a/source/AStack.hpp +++ b/source/AStack.hpp @@ -54,6 +54,10 @@ public: StackType type(); void setName(const std::string &name); size_t numLayers(); + virtual size_t numContext() + { + return 0; + } void addLayer(Layer *pLayer); void delLayer(Layer *pLayer); @@ -67,7 +71,7 @@ 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); + virtual arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch) = 0; size_t numTraining(); void addTraining(const arma::mat &toAdd); diff --git a/source/DeepStack.cpp b/source/DeepStack.cpp index e8337db..08752bf 100644 --- a/source/DeepStack.cpp +++ b/source/DeepStack.cpp @@ -24,6 +24,22 @@ DeepStack::~DeepStack() { } +arma::mat DeepStack::trainingBatchFrom(size_t layerId, const arma::mat& batch) +{ + arma::mat thisBatch = batch; + Layer *pLayer = getLayer(0); + while (pLayer) + { + if (pLayer->id() == layerId) + { + break; + } + thisBatch = pLayer->toHiddenProbs(thisBatch); + pLayer = pLayer->next; + } + return thisBatch; +} + void DeepStack::train(const arma::mat& batch, Rbm::IListener* pListener) { arma::mat thisBatch = batch; diff --git a/source/DeepStack.hpp b/source/DeepStack.hpp index a25be29..9319c2c 100644 --- a/source/DeepStack.hpp +++ b/source/DeepStack.hpp @@ -28,6 +28,7 @@ public: DeepStack(const DeepStack& orig) = delete; virtual ~DeepStack(); + arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch); void train(const arma::mat& batch, Rbm::IListener* pListener) override; void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override; diff --git a/source/Layer.cpp b/source/Layer.cpp index 29a8a63..e4dea72 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -118,23 +118,6 @@ arma::mat Layer::vc_to_v(const arma::mat& vc) const return arma::reshape(vc, 1, numVisible() - m_numContext); } -void Layer::calcContextBatch(arma::mat& batch) -{ - size_t numTraining = batch.n_rows; - - if (m_numContext > 0 and numTraining > 0) - { - for (int i = 1; i < numTraining; i++) - { - arma::mat v = batch.row(i - 1); - arma::mat h = arma::zeros(1, numHidden()); - gibbs_vh(v, h); - arma::mat training_with_ctx = arma::join_rows(batch.row(i).cols(0, numVisible() - m_numContext - 1), h); - batch.row(i) = training_with_ctx; - } - } -} - void Layer::train(const arma::mat& batch, IListener* pListener) { std::cout << m_name << ": " << " Training of layer " << std::to_string(id()) << std::endl; diff --git a/source/Layer.hpp b/source/Layer.hpp index db7d178..ec74cf1 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -48,7 +48,6 @@ public: bool weightsLoad(std::string const &dir, std::string const &prj); bool weightsSave(std::string const &dir, std::string const &prj); - void calcContextBatch(arma::mat &batch); void train(arma::mat const &batch, IListener *pListener=nullptr); arma::mat to_h_gibbs(const arma::mat& v_probs); diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp index 629b248..d106401 100644 --- a/source/MainComponent.cpp +++ b/source/MainComponent.cpp @@ -903,8 +903,7 @@ const juce::String& MainComponent::getBaseDir() void MainComponent::run() { trainButton->setButtonText (TRANS("Stop")); - m_pLayer->calcContextBatch(m_stack->trainingBatch()); -// m_pLayer->train(m_stack->trainingBatch(), this); +// m_pLayer->calcContextBatch(m_stack->trainingBatch()); m_stack->train(m_pLayer->id(), m_stack->trainingBatch(), this); trainButton->setButtonText (TRANS("Train")); } diff --git a/source/RnnStack.cpp b/source/RnnStack.cpp index ffe7ac6..c470cab 100644 --- a/source/RnnStack.cpp +++ b/source/RnnStack.cpp @@ -13,8 +13,9 @@ #include "RnnStack.hpp" -RnnStack::RnnStack(const std::string &name) +RnnStack::RnnStack(const std::string &name, size_t numContext) : AStack(StackType::Rnn, name) +, m_numContext(numContext) { } @@ -22,13 +23,82 @@ RnnStack::~RnnStack() { } +size_t RnnStack::getSeqLen() +{ + return numLayers(); +} + +size_t RnnStack::numContext() +{ + return m_numContext; +} + +arma::mat RnnStack::v_to_vc(const arma::mat& v) const +{ + arma::mat c = arma::zeros(v.n_rows, m_numContext); + return arma::join_rows(v, c); +} + +arma::mat RnnStack::v_to_vc(const arma::mat& v, const arma::mat& c) const +{ + return arma::join_rows(v, c); +} + +arma::mat RnnStack::vc_to_c(const arma::mat& vc) const +{ + size_t numVisible = vc.n_cols; + + if (m_numContext == 0) + { + return arma::mat(1, 0); + } + return vc.submat(0, numVisible - m_numContext, 0, numVisible - 1); +} + +arma::mat RnnStack::vc_to_v(const arma::mat& vc) const +{ + size_t numVisible = vc.n_cols; + return vc.submat(0, 0, vc.n_rows - 1, numVisible - m_numContext - 1); +} + +arma::mat RnnStack::trainingBatchFrom(size_t layerId, const arma::mat& batch) +{ + arma::mat vc = batch; + Layer *pLayer = getLayer(0); + while (pLayer) + { + if (layerId == pLayer->id()) + { + break; + } + arma::mat c = pLayer->to_h_gibbs(vc); + arma::mat v = arma::shift(vc_to_v(vc), layerId, 1); + vc = v_to_vc(v, c); + pLayer = pLayer->next; + } + return vc; +} void RnnStack::train(const arma::mat& batch, Rbm::IListener* pListener) { - + arma::mat thisBatch = batch; + Layer *pLayer = getLayer(0); + while (pLayer) + { + thisBatch = trainingBatchFrom(pLayer->id(), thisBatch); + pLayer->train(thisBatch, pListener); + pLayer = pLayer->next; + } } void RnnStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) { - + arma::mat thisBatch = batch; + Layer *pLayer = getLayer(layerId); + if (pLayer) + { + thisBatch = trainingBatchFrom(pLayer->id(), batch); + pLayer->train(thisBatch, pListener); + } } + diff --git a/source/RnnStack.hpp b/source/RnnStack.hpp index 61d26c0..2d10ee8 100644 --- a/source/RnnStack.hpp +++ b/source/RnnStack.hpp @@ -19,16 +19,24 @@ class RnnStack : public AStack { public: - RnnStack(const std::string &name); + RnnStack(const std::string &name, size_t numContext); RnnStack(const RnnStack& orig) = delete; virtual ~RnnStack(); + arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch) override; void train(const arma::mat& batch, Rbm::IListener* pListener) override; void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override; - + size_t getSeqLen(); + size_t numContext(); + arma::mat v_to_vc(const arma::mat &v) const; + arma::mat v_to_vc(const arma::mat &v, const arma::mat &c) const; + arma::mat vc_to_v(const arma::mat &vc) const; + arma::mat vc_to_c(const arma::mat &vc) const; + private: - + size_t m_numContext; + }; #endif /* RNNSTACK_HPP */ diff --git a/source/StackCreator.cpp b/source/StackCreator.cpp index cf04b57..26c06d5 100644 --- a/source/StackCreator.cpp +++ b/source/StackCreator.cpp @@ -34,11 +34,12 @@ AStack* StackCreator::fromJson(Json::Value& project, LayerConstructor *pLayerCon Json::Value &stack = project["stack"]; Json::Value &layers = project["stack"]["layers"]; int type = stack["type"].asInt(); - + int numContext = stack["numContext"].asInt(); + AStack *pStack = nullptr; if (type == AStack::StackType::Rnn) { - pStack = new RnnStack(name); + pStack = new RnnStack(name, numContext); } else { @@ -94,6 +95,7 @@ bool StackCreator::toFile(AStack* pStack, const std::string& dir, const std::str project["stack"]["name"] = name; project["stack"]["type_string"] = AStack::stackTypeStrings[pStack->type()]; project["stack"]["type"] = pStack->type(); + project["stack"]["numContext"] = pStack->numContext(); Json::Value layers(Json::arrayValue); Layer *pLayer = pStack->getLayer(0); diff --git a/source/poet.cpp b/source/poet.cpp index 41eb5fd..285fe88 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -9,7 +9,7 @@ #include #include "Rbm.hpp" #include "Layer.hpp" -#include "DeepStack.hpp" +#include "RnnStack.hpp" #include "StackCreator.hpp" using namespace std; @@ -140,11 +140,11 @@ arma::mat createTraining(const string &filename, size_t seq_len) struct Rnn { Rnn(Layer *layer) - : nV(layer->numVisible() - layer->context().n_cols) + : nV(layer->numVisible() - layer->numHidden()) , nH(layer->numHidden()) , nVx(layer->numVisibleX()) , nVy(layer->numVisibleY()) - , nC(layer->context().n_cols) + , nC(layer->numHidden()) { } @@ -192,6 +192,7 @@ arma::mat to_next(arma::mat v) } #define CREATE_TRAINING 0 +#define DO_TRAINING 1 int main() { #if CREATE_TRAINING @@ -201,7 +202,7 @@ int main() #endif // Load project - AStack *stack = StackCreator::fromFile(".", "poet5"); + RnnStack *stack = reinterpret_cast(StackCreator::fromFile(".", "poet2")); // Load weights stack->loadWeights("."); @@ -209,6 +210,19 @@ int main() // Load training stack->loadTrainingBatch("."); +#if DO_TRAINING + RbmListener listener; + arma::mat t_vc = stack->trainingBatch(); + + for (int i=0; i < stack->getSeqLen(); i++) + { + stack->getLayer(i)->weightsInit(0.1,0); + stack->train(i, t_vc, &listener); + } + stack->saveWeights("."); + stack->save("."); + +#else Layer *layer = stack->getLayer(0); int numTraining = stack->trainingBatch().n_rows; @@ -268,7 +282,7 @@ int main() putchar(c); } +#endif printf("\n\nEnd of program\n"); - return 0; }