From 8c8e0b324d69f141c1664b4c5c3d347469c0d026 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 10 Jan 2022 12:19:48 +0000 Subject: [PATCH] - assign training data with context - calculate context on demand if empty git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@773 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Layer.cpp | 2 +- source/Layer.hpp | 12 ++++++------ source/MainComponent.cpp | 4 ++-- source/MainComponent.hpp | 9 +++++++++ source/RbmComponent.cpp | 6 +++--- 5 files changed, 21 insertions(+), 12 deletions(-) diff --git a/source/Layer.cpp b/source/Layer.cpp index 0cb8720..c912788 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -23,7 +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) +, 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"); diff --git a/source/Layer.hpp b/source/Layer.hpp index cce7d5b..334303b 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -40,15 +40,15 @@ public: { if (m_numContext) { - arma::mat c_states = arma::zeros(batch.n_rows, m_numContext); + m_context = arma::zeros(batch.n_rows, m_numContext); arma::mat ctx = arma::zeros(1, m_numContext); for (int i=1; i < batch.n_rows; i++) { arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx))); ctx = h; - c_states.row(i) = ctx; + m_context.row(i) = ctx; } - arma::mat batch_with_ctx = arma::join_rows(batch, c_states); + arma::mat batch_with_ctx = arma::join_rows(batch, m_context); Rbm::setBatch(trainingData(batch_with_ctx)); } else @@ -123,11 +123,11 @@ public: return vc.submat(0, numVisible() - m_numContext, 0, numVisible() - 1); } - arma::mat to_vc(const arma::mat &v, const arma::mat &c) const + const arma::mat& context() const { - return arma::join_rows(v, c); + return m_context; } - + private: std::string m_name; diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp index e30b0a1..13fca28 100644 --- a/source/MainComponent.cpp +++ b/source/MainComponent.cpp @@ -682,7 +682,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved) { //[UserSliderCode_patterSlider] -- add your slider handling code here.. m_trainingIndex = (int)sliderThatWasMoved->getValue(); - m_pLayer->setTrainingData(m_stack->trainingData().row(m_trainingIndex)); + m_pLayer->setTrainingData(trainingAt(m_trainingIndex)); //[/UserSliderCode_patterSlider] } else if (sliderThatWasMoved == WeightsSlider) @@ -816,7 +816,7 @@ void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged) m_pLayer->redrawWeights(m_weightIndex); if (m_stack->trainingData().n_rows > 0) { - m_pLayer->setTrainingData(m_stack->trainingData().row(m_trainingIndex)); + m_pLayer->setTrainingData(trainingAt(m_trainingIndex)); } //[/UserComboBoxCode_m_rbmSelect] } diff --git a/source/MainComponent.hpp b/source/MainComponent.hpp index 2f54fd2..e88e4dc 100644 --- a/source/MainComponent.hpp +++ b/source/MainComponent.hpp @@ -124,6 +124,15 @@ private: patterSlider->setRange(0, m_stack->numTraining()-1, 1); } + const arma::mat trainingAt(size_t index) + { + if (m_pLayer->context().is_empty()) + { + m_pLayer->setBatch(m_stack->trainingData()); + } + return arma::join_rows(m_stack->trainingData().row(index), m_pLayer->context().row(index)); + } + void updateControls(); bool onProgress(Rbm *pRbm, const Rbm::Status &status) override; diff --git a/source/RbmComponent.cpp b/source/RbmComponent.cpp index 5ddc3b2..c434fba 100644 --- a/source/RbmComponent.cpp +++ b/source/RbmComponent.cpp @@ -279,12 +279,12 @@ void RbmComponent::gibbs(const arma::mat& vc) arma::mat RbmComponent::getTraining() const { - return to_vc(DrawVisibleTrain->getData(), DrawContextTrain->getData()); + return arma::join_rows(DrawVisibleTrain->getData(), DrawContextTrain->getData()); } arma::mat RbmComponent::getReconst() const { - return to_vc(DrawVisibleReconst->getData(), DrawContextReconst->getData()); + return arma::join_rows(DrawVisibleReconst->getData(), DrawContextReconst->getData()); } void RbmComponent::trainRedraw(const arma::mat& vc) @@ -386,7 +386,7 @@ void RbmComponent::redrawWeights(size_t index) void RbmComponent::setTrainingData(arma::mat const& batch) { RbmComponent *pComp = static_cast (root()); - pComp->upDownPass(to_vc(batch, pComp->DrawContextTrain->getData())); + pComp->upDownPass(batch); } //[/MiscUserCode]