- Rbm::onProgress provides object to itself

- MainComponentz is Rbm::IListener

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@635 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-07 19:27:15 +00:00
parent 8605b37b30
commit 4c68cd4370
4 changed files with 25 additions and 20 deletions
+11 -5
View File
@@ -628,7 +628,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
{ {
//[UserSliderCode_patterSlider] -- add your slider handling code here.. //[UserSliderCode_patterSlider] -- add your slider handling code here..
m_trainingIndex = (int)sliderThatWasMoved->getValue(); m_trainingIndex = (int)sliderThatWasMoved->getValue();
m_pLayer->setTrainingData(m_trainingData.row(m_trainingIndex)); m_pLayer->setTrainingData(m_stack->training().row(m_trainingIndex));
//[/UserSliderCode_patterSlider] //[/UserSliderCode_patterSlider]
} }
else if (sliderThatWasMoved == WeightsSlider) 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... //[MiscUserCode] You can add your own definitions of your custom methods or any other code here...
void MainComponent::save () void MainComponent::save ()
{ {
m_stack->saveTraining(m_trainingData); m_stack->saveTraining();
m_stack->saveWeights(); m_stack->saveWeights();
m_stack->save(); m_stack->save();
} }
@@ -835,12 +835,18 @@ const juce::String& MainComponent::getBaseDir()
void MainComponent::run() 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<RbmComponent*>(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; return !m_doStop;
} }
+10 -11
View File
@@ -42,7 +42,7 @@ class MainComponent
: public Component : public Component
, public Thread , public Thread
, public LayerConstructor , public LayerConstructor
, public RbmComponentListener , public Rbm::IListener
, public ButtonListener , public ButtonListener
, public SliderListener , public SliderListener
, public LabelListener , public LabelListener
@@ -55,7 +55,6 @@ public:
//============================================================================== //==============================================================================
//[UserMethods] -- You can add your own custom methods in this section. //[UserMethods] -- You can add your own custom methods in this section.
bool onProgressChanged(size_t progressPercent) override;
//[/UserMethods] //[/UserMethods]
void paint (Graphics& g); 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) 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); addAndMakeVisible(pComp);
return static_cast<Layer*>(pComp); return static_cast<Layer*>(pComp);
} }
@@ -86,7 +85,6 @@ private:
static const size_t DBN_SIZE = 4; static const size_t DBN_SIZE = 4;
ScopedPointer<Stack> m_stack; ScopedPointer<Stack> m_stack;
RbmComponent *m_pLayer; RbmComponent *m_pLayer;
arma::mat m_trainingData;
int m_weightIndex; int m_weightIndex;
int m_trainingIndex; int m_trainingIndex;
void save(); void save();
@@ -97,13 +95,13 @@ private:
bool m_doStop; bool m_doStop;
void clearTraining() void clearTraining()
{ {
patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); patterSlider->setRange(0, m_stack->numTraining()-1, 1);
} }
void loadTraining() void loadTraining()
{ {
m_trainingData = m_stack->loadTraining(); m_stack->loadTraining();
patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); patterSlider->setRange(0, m_stack->numTraining()-1, 1);
} }
void saveTraining() void saveTraining()
@@ -113,17 +111,18 @@ private:
void addTraining(const arma::mat &training) void addTraining(const arma::mat &training)
{ {
m_stack->addTraining(m_trainingData, training); m_stack->addTraining(training);
patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); patterSlider->setRange(0, m_stack->numTraining()-1, 1);
} }
void removeTrainingAt(size_t index) void removeTrainingAt(size_t index)
{ {
m_stack->delTraining(m_trainingData, index); m_stack->delTraining(index);
patterSlider->setRange(0, m_stack->numTraining(m_trainingData)-1, 1); patterSlider->setRange(0, m_stack->numTraining()-1, 1);
} }
void updateControls(); void updateControls();
bool onProgress(Rbm *pRbm, const Rbm::Status &status) override;
//[/UserVariables] //[/UserVariables]
+3 -3
View File
@@ -81,7 +81,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
if (pListener) if (pListener)
{ {
pListener->onProgress(status); pListener->onProgress(this, status);
} }
while (status.trainingSizeRemain) while (status.trainingSizeRemain)
@@ -181,7 +181,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
if (pListener) if (pListener)
{ {
if(!pListener->onProgress(status)) if(!pListener->onProgress(this, status))
{ {
break; break;
} }
@@ -197,7 +197,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
if (pListener) if (pListener)
{ {
pListener->onProgress(status); pListener->onProgress(this, status);
} }
} }
+1 -1
View File
@@ -112,7 +112,7 @@ public:
IListener() {} IListener() {}
virtual ~IListener() {} virtual ~IListener() {}
virtual bool onProgress(const Status &status) virtual bool onProgress(Rbm *pRbm, const Status &status)
{ {
return true; return true;
} }