- added useProbsForHiddenReconstruction

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@21 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2014-10-08 18:45:15 +00:00
parent da30d95e31
commit fc1f53fdb6
3 changed files with 49 additions and 24 deletions
+33 -16
View File
@@ -195,6 +195,10 @@ MainComponent::MainComponent ()
rbmDoRobbinsMonroToggleButton->setButtonText (TRANS("Robbins-Monro")); rbmDoRobbinsMonroToggleButton->setButtonText (TRANS("Robbins-Monro"));
rbmDoRobbinsMonroToggleButton->addListener (this); rbmDoRobbinsMonroToggleButton->addListener (this);
addAndMakeVisible (rbmReduceVarianceToggleButton = new ToggleButton ("rbmReduceVariance toggle button"));
rbmReduceVarianceToggleButton->setButtonText (TRANS("Reduce Variance"));
rbmReduceVarianceToggleButton->addListener (this);
//[UserPreSize] //[UserPreSize]
m_vNumX = 16; m_vNumX = 16;
@@ -220,6 +224,7 @@ MainComponent::MainComponent ()
rbmUseExpectationsToggleButton->setToggleState(false, sendNotification); rbmUseExpectationsToggleButton->setToggleState(false, sendNotification);
rbmDoRaoBlackwellToggleButton->setToggleState(false, sendNotification); rbmDoRaoBlackwellToggleButton->setToggleState(false, sendNotification);
rbmDoRobbinsMonroToggleButton->setToggleState(false, sendNotification); rbmDoRobbinsMonroToggleButton->setToggleState(false, sendNotification);
rbmReduceVarianceToggleButton->setToggleState(false, sendNotification);
//[/Constructor] //[/Constructor]
} }
@@ -253,6 +258,7 @@ MainComponent::~MainComponent()
rbmUseExpectationsToggleButton = nullptr; rbmUseExpectationsToggleButton = nullptr;
rbmDoRaoBlackwellToggleButton = nullptr; rbmDoRaoBlackwellToggleButton = nullptr;
rbmDoRobbinsMonroToggleButton = nullptr; rbmDoRobbinsMonroToggleButton = nullptr;
rbmReduceVarianceToggleButton = nullptr;
//[Destructor]. You can add your own custom destruction code here.. //[Destructor]. You can add your own custom destruction code here..
@@ -280,14 +286,14 @@ void MainComponent::paint (Graphics& g)
void MainComponent::resized() void MainComponent::resized()
{ {
trainButton->setBounds (120, 312, 72, 24); trainButton->setBounds (120, 312, 72, 24);
addButton->setBounds (24, 176, 72, 24); addButton->setBounds (24, 160, 72, 24);
patterSlider->setBounds (224, 312, 184, 24); patterSlider->setBounds (224, 312, 184, 24);
reconstructButton->setBounds (24, 312, 72, 24); reconstructButton->setBounds (24, 312, 72, 24);
ShakeButton->setBounds (120, 352, 72, 24); ShakeButton->setBounds (120, 352, 72, 24);
WeightsSlider->setBounds (224, 352, 184, 24); WeightsSlider->setBounds (224, 352, 184, 24);
numEpochslabel->setBounds (24, 272, 72, 24); numEpochslabel->setBounds (24, 280, 72, 24);
learningRateLabel->setBounds (120, 272, 72, 24); learningRateLabel->setBounds (120, 280, 72, 24);
testButton->setBounds (120, 176, 72, 24); testButton->setBounds (120, 160, 72, 24);
numVisibleLabel->setBounds (440, 256, 72, 24); numVisibleLabel->setBounds (440, 256, 72, 24);
numHiddenLabel->setBounds (488, 288, 72, 24); numHiddenLabel->setBounds (488, 288, 72, 24);
createButton->setBounds (488, 320, 72, 24); createButton->setBounds (488, 320, 72, 24);
@@ -301,9 +307,10 @@ void MainComponent::resized()
removeTrainingButton->setBounds (440, 88, 72, 24); removeTrainingButton->setBounds (440, 88, 72, 24);
numGibbsSlider->setBounds (224, 272, 184, 24); numGibbsSlider->setBounds (224, 272, 184, 24);
reconstructEquButton->setBounds (24, 352, 72, 24); reconstructEquButton->setBounds (24, 352, 72, 24);
rbmUseExpectationsToggleButton->setBounds (224, 240, 150, 24); rbmUseExpectationsToggleButton->setBounds (24, 232, 128, 24);
rbmDoRaoBlackwellToggleButton->setBounds (224, 176, 150, 24); rbmDoRaoBlackwellToggleButton->setBounds (24, 200, 112, 24);
rbmDoRobbinsMonroToggleButton->setBounds (224, 208, 150, 24); rbmDoRobbinsMonroToggleButton->setBounds (168, 232, 120, 24);
rbmReduceVarianceToggleButton->setBounds (168, 200, 128, 24);
//[UserResized] Add your own custom resize handling here.. //[UserResized] Add your own custom resize handling here..
Draw->setBounds (16, 16, 100, 100); Draw->setBounds (16, 16, 100, 100);
Draw2->setBounds (110+16, 16, 100, 100); Draw2->setBounds (110+16, 16, 100, 100);
@@ -320,7 +327,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
if (buttonThatWasClicked == trainButton) if (buttonThatWasClicked == trainButton)
{ {
//[UserButtonCode_trainButton] -- add your button handler code here.. //[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_rbmDoRobbinsMonro); m_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs, m_rbmUseExpectations, m_rbmDoRaoBlackwell, m_rbmReduceEstimatorVariance, m_rbmDoRobbinsMonro);
//[/UserButtonCode_trainButton] //[/UserButtonCode_trainButton]
} }
else if (buttonThatWasClicked == addButton) else if (buttonThatWasClicked == addButton)
@@ -407,7 +414,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
const double *pV, *pH; const double *pV, *pH;
pV = Draw->getData(); pV = Draw->getData();
for (i=0; i < 1000; i++) for (i=0; i < 100; i++)
{ {
pH = m_pRbm->toHidden(pV); pH = m_pRbm->toHidden(pV);
DrawHidden->setData(pH); DrawHidden->setData(pH);
@@ -434,6 +441,12 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
m_rbmDoRobbinsMonro = buttonThatWasClicked->getToggleState(); m_rbmDoRobbinsMonro = buttonThatWasClicked->getToggleState();
//[/UserButtonCode_rbmDoRobbinsMonroToggleButton] //[/UserButtonCode_rbmDoRobbinsMonroToggleButton]
} }
else if (buttonThatWasClicked == rbmReduceVarianceToggleButton)
{
//[UserButtonCode_rbmReduceVarianceToggleButton] -- add your button handler code here..
m_rbmReduceEstimatorVariance = buttonThatWasClicked->getToggleState();
//[/UserButtonCode_rbmReduceVarianceToggleButton]
}
//[UserbuttonClicked_Post] //[UserbuttonClicked_Post]
//[/UserbuttonClicked_Post] //[/UserbuttonClicked_Post]
@@ -635,7 +648,7 @@ BEGIN_JUCER_METADATA
virtualName="" explicitFocusOrder="0" pos="120 312 72 24" buttonText="Train" virtualName="" explicitFocusOrder="0" pos="120 312 72 24" buttonText="Train"
connectedEdges="0" needsCallback="1" radioGroupId="0"/> connectedEdges="0" needsCallback="1" radioGroupId="0"/>
<TEXTBUTTON name="Add button" id="4609f7b526ec1aef" memberName="addButton" <TEXTBUTTON name="Add button" id="4609f7b526ec1aef" memberName="addButton"
virtualName="" explicitFocusOrder="0" pos="24 176 72 24" buttonText="Add" virtualName="" explicitFocusOrder="0" pos="24 160 72 24" buttonText="Add"
connectedEdges="0" needsCallback="1" radioGroupId="0"/> connectedEdges="0" needsCallback="1" radioGroupId="0"/>
<SLIDER name="Pattern slider" id="c3e0a2c816db81d1" memberName="patterSlider" <SLIDER name="Pattern slider" id="c3e0a2c816db81d1" memberName="patterSlider"
virtualName="" explicitFocusOrder="0" pos="224 312 184 24" min="0" virtualName="" explicitFocusOrder="0" pos="224 312 184 24" min="0"
@@ -652,17 +665,17 @@ BEGIN_JUCER_METADATA
max="0" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft" max="0" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft"
textBoxEditable="1" textBoxWidth="80" textBoxHeight="20" skewFactor="1"/> textBoxEditable="1" textBoxWidth="80" textBoxHeight="20" skewFactor="1"/>
<LABEL name="Num Epochs label" id="b23ae372ee931474" memberName="numEpochslabel" <LABEL name="Num Epochs label" id="b23ae372ee931474" memberName="numEpochslabel"
virtualName="" explicitFocusOrder="0" pos="24 272 72 24" edTextCol="ff000000" virtualName="" explicitFocusOrder="0" pos="24 280 72 24" edTextCol="ff000000"
edBkgCol="0" labelText="99999" editableSingleClick="1" editableDoubleClick="1" edBkgCol="0" labelText="99999" editableSingleClick="1" editableDoubleClick="1"
focusDiscardsChanges="0" fontname="Default font" fontsize="15" focusDiscardsChanges="0" fontname="Default font" fontsize="15"
bold="0" italic="0" justification="36"/> bold="0" italic="0" justification="36"/>
<LABEL name="Learning Rate label" id="49611a27914e910d" memberName="learningRateLabel" <LABEL name="Learning Rate label" id="49611a27914e910d" memberName="learningRateLabel"
virtualName="" explicitFocusOrder="0" pos="120 272 72 24" edTextCol="ff000000" virtualName="" explicitFocusOrder="0" pos="120 280 72 24" edTextCol="ff000000"
edBkgCol="0" labelText="0.001" editableSingleClick="1" editableDoubleClick="1" edBkgCol="0" labelText="0.001" editableSingleClick="1" editableDoubleClick="1"
focusDiscardsChanges="0" fontname="Default font" fontsize="15" focusDiscardsChanges="0" fontname="Default font" fontsize="15"
bold="0" italic="0" justification="36"/> bold="0" italic="0" justification="36"/>
<TEXTBUTTON name="Test button" id="a1e6fed732ae7ee1" memberName="testButton" <TEXTBUTTON name="Test button" id="a1e6fed732ae7ee1" memberName="testButton"
virtualName="" explicitFocusOrder="0" pos="120 176 72 24" buttonText="Test" virtualName="" explicitFocusOrder="0" pos="120 160 72 24" buttonText="Test"
connectedEdges="0" needsCallback="1" radioGroupId="0"/> connectedEdges="0" needsCallback="1" radioGroupId="0"/>
<LABEL name="Num Visible label" id="60acd702770b77dd" memberName="numVisibleLabel" <LABEL name="Num Visible label" id="60acd702770b77dd" memberName="numVisibleLabel"
virtualName="" explicitFocusOrder="0" pos="440 256 72 24" edTextCol="ff000000" virtualName="" explicitFocusOrder="0" pos="440 256 72 24" edTextCol="ff000000"
@@ -714,15 +727,19 @@ BEGIN_JUCER_METADATA
connectedEdges="0" needsCallback="1" radioGroupId="0"/> connectedEdges="0" needsCallback="1" radioGroupId="0"/>
<TOGGLEBUTTON name="rbmUseExpectations toggle button" id="62884e37cd027719" <TOGGLEBUTTON name="rbmUseExpectations toggle button" id="62884e37cd027719"
memberName="rbmUseExpectationsToggleButton" virtualName="" explicitFocusOrder="0" memberName="rbmUseExpectationsToggleButton" virtualName="" explicitFocusOrder="0"
pos="224 240 150 24" buttonText="Use Expectations" connectedEdges="0" pos="24 232 128 24" buttonText="Use Expectations" connectedEdges="0"
needsCallback="1" radioGroupId="0" state="0"/> needsCallback="1" radioGroupId="0" state="0"/>
<TOGGLEBUTTON name="rbmDoRaoBlackwell toggle button" id="7f585de46e182418" <TOGGLEBUTTON name="rbmDoRaoBlackwell toggle button" id="7f585de46e182418"
memberName="rbmDoRaoBlackwellToggleButton" virtualName="" explicitFocusOrder="0" memberName="rbmDoRaoBlackwellToggleButton" virtualName="" explicitFocusOrder="0"
pos="224 176 150 24" buttonText="Rao-Blackwell" connectedEdges="0" pos="24 200 112 24" buttonText="Rao-Blackwell" connectedEdges="0"
needsCallback="1" radioGroupId="0" state="0"/> needsCallback="1" radioGroupId="0" state="0"/>
<TOGGLEBUTTON name="rbmDoRobbinsMonro toggle button" id="21974def5fe90b46" <TOGGLEBUTTON name="rbmDoRobbinsMonro toggle button" id="21974def5fe90b46"
memberName="rbmDoRobbinsMonroToggleButton" virtualName="" explicitFocusOrder="0" memberName="rbmDoRobbinsMonroToggleButton" virtualName="" explicitFocusOrder="0"
pos="224 208 150 24" buttonText="Robbins-Monro" connectedEdges="0" pos="168 232 120 24" buttonText="Robbins-Monro" connectedEdges="0"
needsCallback="1" radioGroupId="0" state="0"/>
<TOGGLEBUTTON name="rbmReduceVariance toggle button" id="1636f6389abe9225"
memberName="rbmReduceVarianceToggleButton" virtualName="" explicitFocusOrder="0"
pos="168 200 128 24" buttonText="Reduce Variance" connectedEdges="0"
needsCallback="1" radioGroupId="0" state="0"/> needsCallback="1" radioGroupId="0" state="0"/>
</JUCER_COMPONENT> </JUCER_COMPONENT>
+2
View File
@@ -91,6 +91,7 @@ private:
bool m_rbmUseExpectations; bool m_rbmUseExpectations;
bool m_rbmDoRaoBlackwell; bool m_rbmDoRaoBlackwell;
bool m_rbmDoRobbinsMonro; bool m_rbmDoRobbinsMonro;
bool m_rbmReduceEstimatorVariance;
//[/UserVariables] //[/UserVariables]
//============================================================================== //==============================================================================
@@ -119,6 +120,7 @@ private:
ScopedPointer<ToggleButton> rbmUseExpectationsToggleButton; ScopedPointer<ToggleButton> rbmUseExpectationsToggleButton;
ScopedPointer<ToggleButton> rbmDoRaoBlackwellToggleButton; ScopedPointer<ToggleButton> rbmDoRaoBlackwellToggleButton;
ScopedPointer<ToggleButton> rbmDoRobbinsMonroToggleButton; ScopedPointer<ToggleButton> rbmDoRobbinsMonroToggleButton;
ScopedPointer<ToggleButton> rbmReduceVarianceToggleButton;
//============================================================================== //==============================================================================
+13 -7
View File
@@ -78,9 +78,9 @@ public:
} }
} }
void train(LayerArray<VisibleLayer> &vt, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations = false, bool doRaoBlackwell = false, bool doRobbinsMonro = false) void train(LayerArray<VisibleLayer> &vt, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations = false, bool doRaoBlackwell = false, bool useProbsForHiddenReconstruction = false, bool doRobbinsMonro = false)
{ {
uint32_t t; uint32_t t, i;
uint32_t epoch; uint32_t epoch;
uint32_t gibbs; uint32_t gibbs;
VisibleLayer v(m_w.getNumVisible()); VisibleLayer v(m_w.getNumVisible());
@@ -100,17 +100,18 @@ public:
if (doRobbinsMonro) if (doRobbinsMonro)
{ {
for (t=0; t < vt.getSize(); t++) for (i=0; i < vt.getSize(); i++)
{ {
// Create hidden layer base on training data // Create hidden layer base on training data
ht[t].probsUpdate(vt[t], w); ht[i].probsUpdate(vt[i], w);
} }
} }
for (epoch=0; epoch < numEpochs; epoch++) for (epoch=0; epoch < numEpochs; epoch++)
{ {
for (t=0; t < vt.getSize(); t++) for (i=0; i < vt.getSize(); i++)
{ {
t = (uint32_t)(0.5 + (vt.getSize()-1)*Noise_Uniform(&m_noise, 0.5));
h.probsUpdate(vt[t], w); h.probsUpdate(vt[t], w);
// Create hidden layer base on training data // Create hidden layer base on training data
@@ -140,7 +141,14 @@ public:
// Create visible reconstruction (a fantasy...) // Create visible reconstruction (a fantasy...)
v.probsUpdate(*pH, w); v.probsUpdate(*pH, w);
if (useProbsForHiddenReconstruction)
{
v.statesAssignfromProbs();
}
else
{
v.statesUpdateStochastic(); v.statesUpdateStochastic();
}
// Create hidden reconstruction // Create hidden reconstruction
pH->probsUpdate(v, w); pH->probsUpdate(v, w);
@@ -284,7 +292,6 @@ public:
m_th.probsUpdate(tv, m_w); m_th.probsUpdate(tv, m_w);
m_th.statesAssignfromProbs(); m_th.statesAssignfromProbs();
// m_th.statesUpdateStochastic();
#if 0 #if 0
printf("pi(t) = (pi^, v>)\n"); printf("pi(t) = (pi^, v>)\n");
@@ -307,7 +314,6 @@ public:
m_tv.probsUpdate(th, m_w); m_tv.probsUpdate(th, m_w);
m_tv.statesAssignfromProbs(); m_tv.statesAssignfromProbs();
// m_tv.statesUpdateStochastic();
#if 0 #if 0
printf("pi(t) = (pi^, v>)\n"); printf("pi(t) = (pi^, v>)\n");