- RBM-Training has own thread

- realtime update of Weights and reconstruction during training
- realtime update of reconstruction on relevant parameters
- added progress bar


git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@26 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2014-10-14 23:06:34 +00:00
parent 7f4438d698
commit 005342e829
4 changed files with 36 additions and 11 deletions
+2 -3
View File
@@ -146,11 +146,10 @@ protected:
void gaussProb(const VectorXd &x, double sigma)
{
uint32_t i;
double var = sigma*sigma;
double k = 1.0/sqrt(2*3.14159265359*var);
double mu = 1;
uint32_t i;
double mu = 0;
for (i=0; i < m_numUnits; i++)
{
+30 -7
View File
@@ -54,7 +54,8 @@ void mylog(const char* format, ...)
//==============================================================================
MainComponent::MainComponent ()
: m_layers(this),
: Thread("RBM"),
m_layers(this),
m_pRbm(nullptr),
DrawTraining(nullptr),
DrawReconstruction(nullptr),
@@ -245,6 +246,9 @@ MainComponent::MainComponent ()
sigmaDecayLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
sigmaDecayLabel->addListener (this);
addAndMakeVisible (m_progressBar = new ProgressBar (m_trainingProgress));
m_progressBar->setName ("Progress Bar");
//[UserPreSize]
m_vNumX = 16;
@@ -263,7 +267,8 @@ MainComponent::MainComponent ()
//[Constructor] You can add your own custom stuff here..
m_trainingProgress = 0;
m_progressBar->setTextToDisplay("Finished");
//[/Constructor]
}
@@ -304,6 +309,7 @@ MainComponent::~MainComponent()
rbmDoSparseToggleButton = nullptr;
sparsityLabel = nullptr;
sigmaDecayLabel = nullptr;
m_progressBar = nullptr;
//[Destructor]. You can add your own custom destruction code here..
@@ -398,6 +404,7 @@ void MainComponent::resized()
rbmDoSparseToggleButton->setBounds (24, 264, 128, 24);
sparsityLabel->setBounds (208, 320, 72, 24);
sigmaDecayLabel->setBounds (304, 320, 72, 24);
m_progressBar->setBounds (224, 160, 168, 24);
//[UserResized] Add your own custom resize handling here..
DrawTraining->setBounds (16, 16, 100, 100);
DrawReconstruction->setBounds (110+16, 16, 100, 100);
@@ -414,7 +421,8 @@ 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_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs);
startThread();
redrawReconstruction();
redrawWeights((int)WeightsSlider->getValue());
//[/UserButtonCode_trainButton]
@@ -478,7 +486,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
m_layers.clear();
m_layers.load((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
redrawReconstruction();
//[/UserButtonCode_loadTrainingButton]
//[/UserButtonCode_loadTrainingButton]
}
else if (buttonThatWasClicked == saveTrainingButton)
{
@@ -730,7 +738,9 @@ void MainComponent::onChanged(const LayerArray<VisibleLayer> &obj)
void MainComponent::onEpochTrained(const Rbm &obj)
{
m_trainingProgress = obj.getProgress();
mylog("Training %f %%\n", m_trainingProgress*100);
// mylog("Training %f %%\n", m_trainingProgress*100);
redrawReconstruction();
redrawWeights((int)WeightsSlider->getValue());
}
void MainComponent::onDraw(DrawComponent &obj)
@@ -776,6 +786,16 @@ void MainComponent::redrawWeights(int index)
DrawWeights->setData(temp);
}
void MainComponent::run()
{
trainButton->setEnabled(false);
m_progressBar->setPercentageDisplay(true);
m_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs);
trainButton->setEnabled(true);
m_trainingProgress = 0.0;
m_progressBar->setTextToDisplay("Finished");
}
//[/MiscUserCode]
@@ -789,8 +809,8 @@ void MainComponent::redrawWeights(int index)
BEGIN_JUCER_METADATA
<JUCER_COMPONENT documentType="Component" className="MainComponent" componentName=""
parentClasses="public Component, public LayerArrayListener&lt;VisibleLayer&gt;, public RbmListener, public DrawListener"
constructorParams="" variableInitialisers="m_layers(this),&#10;m_pRbm(nullptr)&#10;DrawTraining(nullptr),&#10;DrawReconstruction(nullptr),&#10;DrawWeights(nullptr),&#10;DrawHidden(nullptr)"
parentClasses="public Component, public LayerArrayListener&lt;VisibleLayer&gt;, public RbmListener, public DrawListener, public Thread"
constructorParams="" variableInitialisers="Thread(&quot;RBM&quot;),&#10;m_layers(this),&#10;m_pRbm(nullptr)&#10;DrawTraining(nullptr),&#10;DrawReconstruction(nullptr),&#10;DrawWeights(nullptr),&#10;DrawHidden(nullptr)"
snapPixels="8" snapActive="1" snapShown="1" overlayOpacity="0.330"
fixedSize="1" initialWidth="800" initialHeight="600">
<BACKGROUND backgroundColour="ffffffff">
@@ -931,6 +951,9 @@ BEGIN_JUCER_METADATA
edBkgCol="0" labelText="1.0" editableSingleClick="1" editableDoubleClick="1"
focusDiscardsChanges="0" fontname="Default font" fontsize="15"
bold="0" italic="0" justification="36"/>
<GENERICCOMPONENT name="Progress Bar" id="548820d4a30574c7" memberName="m_progressBar"
virtualName="" explicitFocusOrder="0" pos="224 160 168 24" class="ProgressBar"
params="m_trainingProgress"/>
</JUCER_COMPONENT>
END_JUCER_METADATA
+3
View File
@@ -41,6 +41,7 @@ class MainComponent : public Component,
public LayerArrayListener<VisibleLayer>,
public RbmListener,
public DrawListener,
public Thread,
public ButtonListener,
public SliderListener,
public LabelListener
@@ -87,6 +88,7 @@ private:
void onDraw(DrawComponent &obj);
void redrawReconstruction();
void redrawWeights(int index);
void run();
String m_baseDir;
double m_trainingProgress;
uint32_t m_numGibbs;
@@ -125,6 +127,7 @@ private:
ScopedPointer<ToggleButton> rbmDoSparseToggleButton;
ScopedPointer<Label> sparsityLabel;
ScopedPointer<Label> sigmaDecayLabel;
ScopedPointer<ProgressBar> m_progressBar;
//==============================================================================
+1 -1
View File
@@ -210,7 +210,7 @@ public:
m += th.probs();
}
m /= i;
m.array() = m_sparsity - m.array();
m.array() -= m_sparsity;
th.states() = m;
hiddenBiasUpdate(th, -mu);
}