diff --git a/Source/DrawComponent.cpp b/Source/DrawComponent.cpp index 6cbca49..567871c 100644 --- a/Source/DrawComponent.cpp +++ b/Source/DrawComponent.cpp @@ -30,6 +30,7 @@ void mylog(const char* format, ...); //============================================================================== DrawComponent::DrawComponent (int width, int height) + : m_pListener(nullptr) { //[UserPreSize] @@ -130,8 +131,10 @@ void DrawComponent::mouseDown (const MouseEvent& e) } else { - drawAt(e.x, e.y); + drawAt(e.x, e.y, true); } + if (m_pListener) + m_pListener->onDraw(*this); //[/UserCode_mouseDown] } @@ -141,7 +144,9 @@ void DrawComponent::mouseDrag (const MouseEvent& e) // mylog("%s x:%d, y:%d\n", __func__, e.x, e.y); if (e.mods.isLeftButtonDown()) { - drawAt(e.x, e.y); + drawAt(e.x, e.y, false); + if (m_pListener) + m_pListener->onDraw(*this); } //[/UserCode_mouseDrag] } @@ -170,7 +175,12 @@ void DrawComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails //[MiscUserCode] You can add your own definitions of your custom methods or any other code here... -void DrawComponent::drawAt(int x, int y) +void DrawComponent::setListener(DrawListener *pListener) +{ + m_pListener = pListener; +} + +void DrawComponent::drawAt(int x, int y, bool setColor) { int index; double fx = (double)x / m_scaleX; @@ -180,12 +190,29 @@ void DrawComponent::drawAt(int x, int y) // mylog("Write(%d)\n", index); if (index >= 0) + { if (index < m_width*m_height) - m_pData[index] = 1.0; + { + if (setColor) + { + if (m_pData[index] > 0.0) + { + m_currData = 0.0; + m_currColor = (Colours::black); + } + else + { + m_currData = 1.0; + m_currColor = (Colours::white); + } + } + m_pData[index] = m_currData; + m_pG->setColour (m_currColor); + } + } x = (int)fx * m_scaleX; y = (int)fy * m_scaleY; - m_pG->setColour (Colours::white); m_pG->fillRect((float)x, (float)y, m_scaleX, m_scaleY); repaint(); } @@ -239,8 +266,9 @@ BEGIN_JUCER_METADATA + variableInitialisers="m_pListener(nullptr)" snapPixels="8" snapActive="1" + snapShown="1" overlayOpacity="0.330" fixedSize="0" initialWidth="80" + initialHeight="80"> diff --git a/Source/DrawComponent.h b/Source/DrawComponent.h index c37e75d..1c0d6a4 100644 --- a/Source/DrawComponent.h +++ b/Source/DrawComponent.h @@ -22,6 +22,14 @@ //[Headers] -- You can add your own extra header files here -- #include "JuceHeader.h" +class DrawComponent; +class DrawListener +{ +public: + DrawListener() {} + virtual ~DrawListener() {} + virtual void onDraw(DrawComponent &obj) = 0; +}; //[/Headers] @@ -43,7 +51,8 @@ public: //============================================================================== //[UserMethods] -- You can add your own custom methods in this section. - void drawAt(int x, int y); + void setListener(DrawListener *pListener); + void drawAt(int x, int y, bool setColor); void setData(const double *pData); const double* getData(); void clear(); @@ -64,6 +73,7 @@ public: private: //[UserVariables] -- You can add your own custom variables in this section. + DrawListener *m_pListener; int m_width; int m_height; float m_scaleX; @@ -71,6 +81,8 @@ private: ScopedPointerm_pG; ScopedPointerm_pData; Image m_image; + double m_currData; + Colour m_currColor; //[/UserVariables] //============================================================================== diff --git a/Source/LayerArray.hpp b/Source/LayerArray.hpp index edf65d7..9e899a5 100644 --- a/Source/LayerArray.hpp +++ b/Source/LayerArray.hpp @@ -55,6 +55,20 @@ public: { } + LayerArray(uint32_t numLayers, uint32_t numUnitsPerLayer) + : m_size(0) + , m_pRoot(nullptr) + , m_ppIndex(nullptr) + , m_pListener(nullptr) + { + uint32_t i; + + for (i=0; i < numLayers; i++) + { + add(nullptr, numUnitsPerLayer); + } + } + virtual ~LayerArray() { m_pListener = nullptr; diff --git a/Source/MainComponent.cpp b/Source/MainComponent.cpp index ee6c49b..1655b40 100644 --- a/Source/MainComponent.cpp +++ b/Source/MainComponent.cpp @@ -56,7 +56,8 @@ MainComponent::MainComponent () m_pRbm(nullptr), Draw(nullptr), Draw2(nullptr), - DrawWeights(nullptr) + DrawWeights(nullptr), + DrawHidden(nullptr) { addAndMakeVisible (trainButton = new TextButton ("Train button")); trainButton->setButtonText (TRANS("Train")); @@ -190,9 +191,9 @@ MainComponent::MainComponent () rbmDoRaoBlackwellToggleButton->setButtonText (TRANS("Rao-Blackwell")); rbmDoRaoBlackwellToggleButton->addListener (this); - addAndMakeVisible (rbmDoRobinsMonroToggleButton = new ToggleButton ("rbmDoRobinsMonro toggle button")); - rbmDoRobinsMonroToggleButton->setButtonText (TRANS("Robins-Monro")); - rbmDoRobinsMonroToggleButton->addListener (this); + addAndMakeVisible (rbmDoRobbinsMonroToggleButton = new ToggleButton ("rbmDoRobbinsMonro toggle button")); + rbmDoRobbinsMonroToggleButton->setButtonText (TRANS("Robbins-Monro")); + rbmDoRobbinsMonroToggleButton->addListener (this); //[UserPreSize] @@ -216,9 +217,9 @@ MainComponent::MainComponent () projectNameLabel->setText(String("TestPrj"), dontSendNotification ); numEpochslabel->setText(String(100), dontSendNotification ); learningRateLabel->setText(String(0.2), dontSendNotification ); - rbmUseExpectationsToggleButton->setToggleState(false, true); - rbmDoRaoBlackwellToggleButton->setToggleState(false, true); - rbmDoRobinsMonroToggleButton->setToggleState(false, true); + rbmUseExpectationsToggleButton->setToggleState(false, sendNotification); + rbmDoRaoBlackwellToggleButton->setToggleState(false, sendNotification); + rbmDoRobbinsMonroToggleButton->setToggleState(false, sendNotification); //[/Constructor] } @@ -251,13 +252,14 @@ MainComponent::~MainComponent() reconstructEquButton = nullptr; rbmUseExpectationsToggleButton = nullptr; rbmDoRaoBlackwellToggleButton = nullptr; - rbmDoRobinsMonroToggleButton = nullptr; + rbmDoRobbinsMonroToggleButton = nullptr; //[Destructor]. You can add your own custom destruction code here.. Draw = nullptr; Draw2 = nullptr; DrawWeights = nullptr; + DrawHidden = nullptr; m_pRbm = nullptr; //[/Destructor] @@ -301,11 +303,12 @@ void MainComponent::resized() reconstructEquButton->setBounds (24, 352, 72, 24); rbmUseExpectationsToggleButton->setBounds (224, 240, 150, 24); rbmDoRaoBlackwellToggleButton->setBounds (224, 176, 150, 24); - rbmDoRobinsMonroToggleButton->setBounds (224, 208, 150, 24); + rbmDoRobbinsMonroToggleButton->setBounds (224, 208, 150, 24); //[UserResized] Add your own custom resize handling here.. Draw->setBounds (16, 16, 100, 100); Draw2->setBounds (110+16, 16, 100, 100); DrawWeights->setBounds (220+16, 16, 100, 100); + DrawHidden->setBounds (16, 16+110, 320, 20); //[/UserResized] } @@ -317,7 +320,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) if (buttonThatWasClicked == trainButton) { //[UserButtonCode_trainButton] -- add your button handler code here.. - m_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs, m_rbmUseExpectations, m_rbmDoRaoBlackwell, m_rbmDoRobinsMonro); + m_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs, m_rbmUseExpectations, m_rbmDoRaoBlackwell, m_rbmDoRobbinsMonro); //[/UserButtonCode_trainButton] } else if (buttonThatWasClicked == addButton) @@ -332,6 +335,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) { //[UserButtonCode_reconstructButton] -- add your button handler code here.. Draw2->setData(m_pRbm->toVisible(m_pRbm->toHidden(Draw->getData()))); + DrawHidden->setData(m_pRbm->toHidden(Draw->getData())); //[/UserButtonCode_reconstructButton] } else if (buttonThatWasClicked == ShakeButton) @@ -406,6 +410,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) for (i=0; i < 1000; i++) { pH = m_pRbm->toHidden(pV); + DrawHidden->setData(pH); pV = m_pRbm->toVisible(pH); Draw2->setData(pV); } @@ -423,11 +428,11 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) m_rbmDoRaoBlackwell = buttonThatWasClicked->getToggleState(); //[/UserButtonCode_rbmDoRaoBlackwellToggleButton] } - else if (buttonThatWasClicked == rbmDoRobinsMonroToggleButton) + else if (buttonThatWasClicked == rbmDoRobbinsMonroToggleButton) { - //[UserButtonCode_rbmDoRobinsMonroToggleButton] -- add your button handler code here.. - m_rbmDoRobinsMonro = buttonThatWasClicked->getToggleState(); - //[/UserButtonCode_rbmDoRobinsMonroToggleButton] + //[UserButtonCode_rbmDoRobbinsMonroToggleButton] -- add your button handler code here.. + m_rbmDoRobbinsMonro = buttonThatWasClicked->getToggleState(); + //[/UserButtonCode_rbmDoRobbinsMonroToggleButton] } //[UserbuttonClicked_Post] @@ -556,10 +561,14 @@ void MainComponent::create() Draw = nullptr; Draw2 = nullptr; DrawWeights = nullptr; + DrawHidden = nullptr; m_pRbm = nullptr; - addAndMakeVisible (Draw = new DrawComponent (m_vNumX, m_vNumX)); - addAndMakeVisible (Draw2 = new DrawComponent (m_vNumX, m_vNumX)); - addAndMakeVisible (DrawWeights = new DrawComponent (m_vNumX, m_vNumX)); + addAndMakeVisible (Draw = new DrawComponent (m_vNumX, m_vNumY)); + Draw->setListener(this); + addAndMakeVisible (Draw2 = new DrawComponent (m_vNumX, m_vNumY)); + addAndMakeVisible (DrawWeights = new DrawComponent (m_vNumX, m_vNumY)); + addAndMakeVisible (DrawHidden = new DrawComponent (m_hNum, 1)); + DrawHidden->setListener(this); numVisibleLabel->setText(String(m_vNumX), dontSendNotification ); numVisibleYLabel->setText(String(m_vNumY), dontSendNotification ); numHiddenLabel->setText(String(m_hNum), dontSendNotification ); @@ -591,6 +600,19 @@ void MainComponent::onEpochTrained(const Rbm &obj) mylog("Training %f %%\n", m_trainingProgress*100); } +void MainComponent::onDraw(DrawComponent &obj) +{ + if (&obj == DrawHidden) + { + Draw2->setData(m_pRbm->toVisible(obj.getData())); + } + if (&obj == Draw) + { + Draw2->setData(m_pRbm->toVisible(m_pRbm->toHidden(obj.getData()))); + DrawHidden->setData(m_pRbm->toHidden(obj.getData())); + } +} + //[/MiscUserCode] @@ -604,8 +626,8 @@ void MainComponent::onEpochTrained(const Rbm &obj) BEGIN_JUCER_METADATA @@ -698,9 +720,10 @@ BEGIN_JUCER_METADATA memberName="rbmDoRaoBlackwellToggleButton" virtualName="" explicitFocusOrder="0" pos="224 176 150 24" buttonText="Rao-Blackwell" connectedEdges="0" needsCallback="1" radioGroupId="0" state="0"/> - + END_JUCER_METADATA diff --git a/Source/MainComponent.h b/Source/MainComponent.h index 4169f62..e38cb52 100644 --- a/Source/MainComponent.h +++ b/Source/MainComponent.h @@ -40,6 +40,7 @@ class MainComponent : public Component, public LayerArrayListener, public RbmListener, + public DrawListener, public ButtonListener, public SliderListener, public LabelListener @@ -68,6 +69,7 @@ private: ScopedPointer Draw; ScopedPointer Draw2; ScopedPointer DrawWeights; + ScopedPointer DrawHidden; uint32_t m_vNumX; uint32_t m_vNumY; uint32_t m_hNum; @@ -82,12 +84,13 @@ private: const juce::String& getBaseDir(); void onChanged(const LayerArray &obj); void onEpochTrained(const Rbm &obj); + void onDraw(DrawComponent &obj); String m_baseDir; double m_trainingProgress; uint32_t m_numGibbs; bool m_rbmUseExpectations; bool m_rbmDoRaoBlackwell; - bool m_rbmDoRobinsMonro; + bool m_rbmDoRobbinsMonro; //[/UserVariables] //============================================================================== @@ -115,7 +118,7 @@ private: ScopedPointer reconstructEquButton; ScopedPointer rbmUseExpectationsToggleButton; ScopedPointer rbmDoRaoBlackwellToggleButton; - ScopedPointer rbmDoRobinsMonroToggleButton; + ScopedPointer rbmDoRobbinsMonroToggleButton; //============================================================================== diff --git a/Source/Rbm.hpp b/Source/Rbm.hpp index 9e9b50f..3db3a21 100644 --- a/Source/Rbm.hpp +++ b/Source/Rbm.hpp @@ -78,57 +78,83 @@ public: } } - void train(LayerArray &vts, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations = false, bool doRaoBlackwell = false, bool doRobinsMonro = false) + void train(LayerArray &vt, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations = false, bool doRaoBlackwell = false, bool doRobbinsMonro = false) { uint32_t t; uint32_t epoch; uint32_t gibbs; VisibleLayer v(m_w.getNumVisible()); HiddenLayer h(m_w.getNumHidden()); + HiddenLayer *pH; + LayerArray ht(vt.getSize(), m_w.getNumHidden()); + Weights w = m_w; double dProgress = 1.0/numEpochs; m_progress = 0; - for (epoch=0; epoch < numEpochs; epoch++) + if (useExpectations) { - for (t=0; t < vts.getSize(); t++) + mu /= vt.getSize(); + } + + if (doRobbinsMonro) + { + for (t=0; t < vt.getSize(); t++) { // Create hidden layer base on training data - h.probsUpdate(vts[t], w); + ht[t].probsUpdate(vt[t], w); + } + } + + for (epoch=0; epoch < numEpochs; epoch++) + { + for (t=0; t < vt.getSize(); t++) + { + h.probsUpdate(vt[t], w); + + // Create hidden layer base on training data + if (doRobbinsMonro) + { + pH = &ht[t]; + } + else + { + pH = &h; + } // Update weights (positive phase) if (doRaoBlackwell) { - h.statesAssignfromProbs(); + pH->statesAssignfromProbs(); } else { - h.statesUpdateStochastic(); + pH->statesUpdateStochastic(); } - weightsUpdate(vts[t], h, +mu/vts.getSize()); + weightsUpdate(vt[t], h, +mu); for (gibbs=0; gibbs < numGibbs; gibbs++) { - h.statesUpdateStochastic(); + pH->statesUpdateStochastic(); // Create visible reconstruction (a fantasy...) - v.probsUpdate(h, w); + v.probsUpdate(*pH, w); v.statesUpdateStochastic(); // Create hidden reconstruction - h.probsUpdate(v, w); + pH->probsUpdate(v, w); } // Update weights (negative phase) if (doRaoBlackwell) { - h.statesAssignfromProbs(); + pH->statesAssignfromProbs(); } else { - h.statesUpdateStochastic(); + pH->statesUpdateStochastic(); } - weightsUpdate(v, h, -mu/vts.getSize()); + weightsUpdate(v, *pH, -mu); if (!useExpectations) {