- 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:
@@ -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
@@ -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
@@ -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
@@ -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;
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user