diff --git a/Source/MainComponent.cpp b/Source/MainComponent.cpp
index 2a77744..ee6c49b 100644
--- a/Source/MainComponent.cpp
+++ b/Source/MainComponent.cpp
@@ -186,6 +186,14 @@ MainComponent::MainComponent ()
rbmUseExpectationsToggleButton->setButtonText (TRANS("Use Expectations"));
rbmUseExpectationsToggleButton->addListener (this);
+ addAndMakeVisible (rbmDoRaoBlackwellToggleButton = new ToggleButton ("rbmDoRaoBlackwell toggle button"));
+ rbmDoRaoBlackwellToggleButton->setButtonText (TRANS("Rao-Blackwell"));
+ rbmDoRaoBlackwellToggleButton->addListener (this);
+
+ addAndMakeVisible (rbmDoRobinsMonroToggleButton = new ToggleButton ("rbmDoRobinsMonro toggle button"));
+ rbmDoRobinsMonroToggleButton->setButtonText (TRANS("Robins-Monro"));
+ rbmDoRobinsMonroToggleButton->addListener (this);
+
//[UserPreSize]
m_vNumX = 16;
@@ -209,6 +217,8 @@ MainComponent::MainComponent ()
numEpochslabel->setText(String(100), dontSendNotification );
learningRateLabel->setText(String(0.2), dontSendNotification );
rbmUseExpectationsToggleButton->setToggleState(false, true);
+ rbmDoRaoBlackwellToggleButton->setToggleState(false, true);
+ rbmDoRobinsMonroToggleButton->setToggleState(false, true);
//[/Constructor]
}
@@ -240,6 +250,8 @@ MainComponent::~MainComponent()
numGibbsSlider = nullptr;
reconstructEquButton = nullptr;
rbmUseExpectationsToggleButton = nullptr;
+ rbmDoRaoBlackwellToggleButton = nullptr;
+ rbmDoRobinsMonroToggleButton = nullptr;
//[Destructor]. You can add your own custom destruction code here..
@@ -265,14 +277,14 @@ void MainComponent::paint (Graphics& g)
void MainComponent::resized()
{
- trainButton->setBounds (120, 280, 72, 24);
+ trainButton->setBounds (120, 312, 72, 24);
addButton->setBounds (24, 176, 72, 24);
- patterSlider->setBounds (224, 280, 184, 24);
- reconstructButton->setBounds (24, 280, 72, 24);
- ShakeButton->setBounds (120, 320, 72, 24);
- WeightsSlider->setBounds (224, 320, 184, 24);
- numEpochslabel->setBounds (24, 240, 72, 24);
- learningRateLabel->setBounds (120, 240, 72, 24);
+ patterSlider->setBounds (224, 312, 184, 24);
+ reconstructButton->setBounds (24, 312, 72, 24);
+ ShakeButton->setBounds (120, 352, 72, 24);
+ WeightsSlider->setBounds (224, 352, 184, 24);
+ numEpochslabel->setBounds (24, 272, 72, 24);
+ learningRateLabel->setBounds (120, 272, 72, 24);
testButton->setBounds (120, 176, 72, 24);
numVisibleLabel->setBounds (440, 256, 72, 24);
numHiddenLabel->setBounds (488, 288, 72, 24);
@@ -285,9 +297,11 @@ void MainComponent::resized()
saveTrainingButton->setBounds (520, 48, 72, 24);
clearTrainingButton->setBounds (520, 88, 72, 24);
removeTrainingButton->setBounds (440, 88, 72, 24);
- numGibbsSlider->setBounds (224, 240, 184, 24);
- reconstructEquButton->setBounds (24, 320, 72, 24);
- rbmUseExpectationsToggleButton->setBounds (224, 176, 150, 24);
+ numGibbsSlider->setBounds (224, 272, 184, 24);
+ reconstructEquButton->setBounds (24, 352, 72, 24);
+ rbmUseExpectationsToggleButton->setBounds (224, 240, 150, 24);
+ rbmDoRaoBlackwellToggleButton->setBounds (224, 176, 150, 24);
+ rbmDoRobinsMonroToggleButton->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);
@@ -303,7 +317,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_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs, m_rbmUseExpectations, m_rbmDoRaoBlackwell, m_rbmDoRobinsMonro);
//[/UserButtonCode_trainButton]
}
else if (buttonThatWasClicked == addButton)
@@ -403,6 +417,18 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
m_rbmUseExpectations = buttonThatWasClicked->getToggleState();
//[/UserButtonCode_rbmUseExpectationsToggleButton]
}
+ else if (buttonThatWasClicked == rbmDoRaoBlackwellToggleButton)
+ {
+ //[UserButtonCode_rbmDoRaoBlackwellToggleButton] -- add your button handler code here..
+ m_rbmDoRaoBlackwell = buttonThatWasClicked->getToggleState();
+ //[/UserButtonCode_rbmDoRaoBlackwellToggleButton]
+ }
+ else if (buttonThatWasClicked == rbmDoRobinsMonroToggleButton)
+ {
+ //[UserButtonCode_rbmDoRobinsMonroToggleButton] -- add your button handler code here..
+ m_rbmDoRobinsMonro = buttonThatWasClicked->getToggleState();
+ //[/UserButtonCode_rbmDoRobinsMonroToggleButton]
+ }
//[UserbuttonClicked_Post]
//[/UserbuttonClicked_Post]
@@ -584,32 +610,32 @@ BEGIN_JUCER_METADATA
fixedSize="1" initialWidth="600" initialHeight="400">
@@ -658,16 +684,23 @@ BEGIN_JUCER_METADATA
virtualName="" explicitFocusOrder="0" pos="440 88 72 24" buttonText="Remove T"
connectedEdges="0" needsCallback="1" radioGroupId="0"/>
+
+
END_JUCER_METADATA
diff --git a/Source/MainComponent.h b/Source/MainComponent.h
index 03cfe96..4169f62 100644
--- a/Source/MainComponent.h
+++ b/Source/MainComponent.h
@@ -86,6 +86,8 @@ private:
double m_trainingProgress;
uint32_t m_numGibbs;
bool m_rbmUseExpectations;
+ bool m_rbmDoRaoBlackwell;
+ bool m_rbmDoRobinsMonro;
//[/UserVariables]
//==============================================================================
@@ -112,6 +114,8 @@ private:
ScopedPointer numGibbsSlider;
ScopedPointer reconstructEquButton;
ScopedPointer rbmUseExpectationsToggleButton;
+ ScopedPointer rbmDoRaoBlackwellToggleButton;
+ ScopedPointer rbmDoRobinsMonroToggleButton;
//==============================================================================
diff --git a/Source/Rbm.hpp b/Source/Rbm.hpp
index 0d0b20b..9e9b50f 100644
--- a/Source/Rbm.hpp
+++ b/Source/Rbm.hpp
@@ -78,7 +78,7 @@ public:
}
}
- void train(LayerArray &vts, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations=false)
+ void train(LayerArray &vts, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations = false, bool doRaoBlackwell = false, bool doRobinsMonro = false)
{
uint32_t t;
uint32_t epoch;
@@ -96,23 +96,38 @@ public:
{
// Create hidden layer base on training data
h.probsUpdate(vts[t], w);
- h.statesUpdateStochastic();
// Update weights (positive phase)
+ if (doRaoBlackwell)
+ {
+ h.statesAssignfromProbs();
+ }
+ else
+ {
+ h.statesUpdateStochastic();
+ }
weightsUpdate(vts[t], h, +mu/vts.getSize());
for (gibbs=0; gibbs < numGibbs; gibbs++)
{
+ h.statesUpdateStochastic();
+
// Create visible reconstruction (a fantasy...)
v.probsUpdate(h, w);
v.statesUpdateStochastic();
// Create hidden reconstruction
h.probsUpdate(v, w);
- h.statesUpdateStochastic();
}
// Update weights (negative phase)
- h.statesAssignfromProbs();
+ if (doRaoBlackwell)
+ {
+ h.statesAssignfromProbs();
+ }
+ else
+ {
+ h.statesUpdateStochastic();
+ }
weightsUpdate(v, h, -mu/vts.getSize());
if (!useExpectations)