From 648d8a002ea29c7af7fe65d3c4e90c30a5f11b23 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 10 Jan 2022 10:36:55 +0000 Subject: [PATCH] - added setBatch() git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@771 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Layer.cpp | 1 + source/Layer.hpp | 9 +++++---- source/MainComponent.cpp | 3 ++- source/Rbm.cpp | 3 ++- source/Rbm.hpp | 8 +++++++- source/Stack.cpp | 3 ++- 6 files changed, 19 insertions(+), 8 deletions(-) diff --git a/source/Layer.cpp b/source/Layer.cpp index 3ca066f..5f8731b 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -23,6 +23,7 @@ Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibl , m_numVisibleX(numVisibleX) , m_numVisibleY(numVisibleY) , m_numContext(numContext) +, m_context(1, numContext) { cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl; m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat"); diff --git a/source/Layer.hpp b/source/Layer.hpp index 69ffe90..cce7d5b 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -33,8 +33,8 @@ public: Json::Value toJson() const; bool loadWeights(const std::string &prjname=""); bool saveWeights(const std::string &prjname=""); - - void train(arma::mat const &batch, IListener *pListener=nullptr) + + void setBatch(arma::mat const &batch) { if (batch.n_rows > 0) { @@ -49,11 +49,11 @@ public: c_states.row(i) = ctx; } arma::mat batch_with_ctx = arma::join_rows(batch, c_states); - Rbm::train(trainingData(batch_with_ctx), pListener); + Rbm::setBatch(trainingData(batch_with_ctx)); } else { - Rbm::train(trainingData(batch), pListener); + Rbm::setBatch(trainingData(batch)); } } } @@ -136,6 +136,7 @@ private: size_t m_numVisibleX; size_t m_numVisibleY; size_t m_numContext; + arma::mat m_context; }; diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp index 15d02d8..e30b0a1 100644 --- a/source/MainComponent.cpp +++ b/source/MainComponent.cpp @@ -899,7 +899,8 @@ const juce::String& MainComponent::getBaseDir() void MainComponent::run() { trainButton->setButtonText (TRANS("Stop")); - m_pLayer->train(m_stack->trainingData(), this); + m_pLayer->setBatch(m_stack->trainingData()); + m_pLayer->train(this); trainButton->setButtonText (TRANS("Train")); } diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 3a7f676..3af22aa 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -116,9 +116,10 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma:: dbhv -= sum(h_probs, 0); } -void Rbm::train(const arma::mat& batch, IListener* pListener) +void Rbm::train(IListener* pListener) { Status status; + arma::mat const &batch = m_batch; double dProgress = 100.0/(batch.n_rows*m_params.numEpochs); double progress = 0; int lastProgress = -100; diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 163f971..fb8390b 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -126,7 +126,12 @@ public: m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv; } - void train(arma::mat const &batch, IListener *pListener=nullptr); + void setBatch(arma::mat const &batch) + { + m_batch = batch; + } + + void train(IListener *pListener=nullptr); static arma::mat normalize(const arma::mat &hidden); const arma::mat& whv() const; @@ -179,6 +184,7 @@ protected: private: noise_gen_t m_noise; arma::mat m_whv; + arma::mat m_batch; }; diff --git a/source/Stack.cpp b/source/Stack.cpp index b1f280c..9ed3358 100644 --- a/source/Stack.cpp +++ b/source/Stack.cpp @@ -204,7 +204,8 @@ void Stack::train(Rbm::IListener* pListener) while(pLayer) { std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl; - pLayer->train(m_trainingData, pListener); + pLayer->setBatch(m_trainingData); + pLayer->train(pListener); pLayer = pLayer->next; } }