diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp index 2d000fe..a9de2c9 100644 --- a/source/MainComponent.cpp +++ b/source/MainComponent.cpp @@ -628,7 +628,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved) { //[UserSliderCode_patterSlider] -- add your slider handling code here.. m_trainingIndex = (int)sliderThatWasMoved->getValue(); - m_pLayer->setTrainingData(m_trainingData.row(m_trainingIndex)); + m_pLayer->setTrainingData(m_stack->training().row(m_trainingIndex)); //[/UserSliderCode_patterSlider] } else if (sliderThatWasMoved == WeightsSlider) @@ -815,7 +815,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_trainingData); + m_stack->saveTraining(); m_stack->saveWeights(); m_stack->save(); } @@ -835,12 +835,18 @@ const juce::String& MainComponent::getBaseDir() void MainComponent::run() { - m_pLayer->train(m_trainingData); + m_stack->train(this); } -bool MainComponent::onProgressChanged(size_t progressPercent) +bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status) { - m_progressBarSlider->setValue(progressPercent); + RbmComponent *pComp = static_cast(pRbm); + std::cout << "Progress changed of " << pComp->name() << "." << std::to_string((int)pComp->id()) << std::endl; + pComp->upPass(pComp->DrawHidden->getData()); + pComp->redrawReconstruction(); + pComp->redrawWeights(); + + m_progressBarSlider->setValue((int)(100*status.progress + 0.5)); return !m_doStop; } diff --git a/source/MainComponent.hpp b/source/MainComponent.hpp index 6379db9..229c4d9 100644 --- a/source/MainComponent.hpp +++ b/source/MainComponent.hpp @@ -42,7 +42,7 @@ class MainComponent : public Component , public Thread , public LayerConstructor -, public RbmComponentListener +, public Rbm::IListener , public ButtonListener , public SliderListener , public LabelListener @@ -55,7 +55,6 @@ public: //============================================================================== //[UserMethods] -- You can add your own custom methods in this section. - bool onProgressChanged(size_t progressPercent) override; //[/UserMethods] void paint (Graphics& g); @@ -75,7 +74,7 @@ public: Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) { - RbmComponent *pComp = new RbmComponent(name, id, numVisibleX, numVisibleY, numHidden, this); + RbmComponent *pComp = new RbmComponent(name, id, numVisibleX, numVisibleY, numHidden); addAndMakeVisible(pComp); return static_cast(pComp); } @@ -86,7 +85,6 @@ private: static const size_t DBN_SIZE = 4; ScopedPointer m_stack; RbmComponent *m_pLayer; - arma::mat m_trainingData; int m_weightIndex; int m_trainingIndex; void save(); @@ -97,13 +95,13 @@ private: bool m_doStop; void clearTraining() { - patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); + patterSlider->setRange(0, m_stack->numTraining()-1, 1); } void loadTraining() { - m_trainingData = m_stack->loadTraining(); - patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); + m_stack->loadTraining(); + patterSlider->setRange(0, m_stack->numTraining()-1, 1); } void saveTraining() @@ -113,17 +111,18 @@ private: void addTraining(const arma::mat &training) { - m_stack->addTraining(m_trainingData, training); - patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); + m_stack->addTraining(training); + patterSlider->setRange(0, m_stack->numTraining()-1, 1); } void removeTrainingAt(size_t index) { - m_stack->delTraining(m_trainingData, index); - patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); + m_stack->delTraining(index); + patterSlider->setRange(0, m_stack->numTraining()-1, 1); } void updateControls(); + bool onProgress(Rbm *pRbm, const Rbm::Status &status) override; //[/UserVariables] diff --git a/source/Rbm.cpp b/source/Rbm.cpp index f8185f8..dfb7166 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -81,7 +81,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) if (pListener) { - pListener->onProgress(status); + pListener->onProgress(this, status); } while (status.trainingSizeRemain) @@ -181,7 +181,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) if (pListener) { - if(!pListener->onProgress(status)) + if(!pListener->onProgress(this, status)) { break; } @@ -197,7 +197,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) if (pListener) { - pListener->onProgress(status); + pListener->onProgress(this, status); } } diff --git a/source/Rbm.hpp b/source/Rbm.hpp index a08e49b..7463bfb 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -112,7 +112,7 @@ public: IListener() {} virtual ~IListener() {} - virtual bool onProgress(const Status &status) + virtual bool onProgress(Rbm *pRbm, const Status &status) { return true; }