From 46bffda38f112ab779eb4fc12a2ddd834a0a4a0c Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 17 Jun 2016 22:14:55 +0000 Subject: [PATCH] [RBM] - added DBN stack - concentrated RBM params into structure git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@295 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- Source/MainComponent.cpp | 215 +++++++++++++++++++++++--------------- Source/MainComponent.h | 22 ++-- Source/Rbm.hpp | 220 ++++++++++++++++++++++----------------- Source/RbmComponent.cpp | 36 +++++-- Source/RbmComponent.h | 24 ++++- 5 files changed, 319 insertions(+), 198 deletions(-) diff --git a/Source/MainComponent.cpp b/Source/MainComponent.cpp index 73d48d4..fd5b4d8 100644 --- a/Source/MainComponent.cpp +++ b/Source/MainComponent.cpp @@ -38,10 +38,9 @@ void mylog(const char* format, ...) //============================================================================== MainComponent::MainComponent () : Thread("RBM"), - m_pRbmComponent(nullptr), - m_layers(this) - - + m_pRbmComponentCurr(nullptr), + m_weightsCurr(nullptr), + m_layers(this) { addAndMakeVisible (trainButton = new TextButton ("Train button")); trainButton->setButtonText (TRANS("Train")); @@ -269,10 +268,20 @@ MainComponent::MainComponent () rbmNormalizeDataToggleButton->setButtonText (TRANS("Normalize data")); rbmNormalizeDataToggleButton->addListener (this); + addAndMakeVisible (m_rbmSelect = new ComboBox ("RBM Selector")); + m_rbmSelect->setEditableText (false); + m_rbmSelect->setJustificationType (Justification::centredLeft); + m_rbmSelect->setTextWhenNothingSelected (String::empty); + m_rbmSelect->setTextWhenNoChoicesAvailable (TRANS("(no choices)")); + m_rbmSelect->addListener (this); + //[UserPreSize] + memset(m_pRbmComponent, 0, sizeof(m_pRbmComponent)); + memset(m_weights, 0, sizeof(m_weights)); + m_rbmSelect->addItem(String(0), 1); + m_rbmSelect->setSelectedId(1, dontSendNotification); m_progressBarSlider->setValue(100); - create(); //[/UserPreSize] setSize (1000, 600); @@ -324,10 +333,10 @@ MainComponent::~MainComponent() weightInitLabel = nullptr; rbmLearnVarianceButton = nullptr; rbmNormalizeDataToggleButton = nullptr; + m_rbmSelect = nullptr; //[Destructor]. You can add your own custom destruction code here.. - m_pRbmComponent = nullptr; //[/Destructor] } @@ -443,6 +452,7 @@ void MainComponent::resized() weightInitLabel->setBounds (888, 200, 72, 24); rbmLearnVarianceButton->setBounds (808, 60, 128, 24); rbmNormalizeDataToggleButton->setBounds (808, 92, 160, 24); + m_rbmSelect->setBounds (844, 404, 62, 24); //[UserResized] Add your own custom resize handling here.. //[/UserResized] } @@ -461,29 +471,30 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) else if (buttonThatWasClicked == addButton) { //[UserButtonCode_addButton] -- add your button handler code here.. - m_layers.add(m_pRbmComponent->getTrainingData(), m_weights->getNumVisible()); - m_pRbmComponent->addFromTraining(); + m_layers.add(m_pRbmComponent[0]->getTrainingData(), m_weights[0]->getNumVisible()); + m_pRbmComponent[0]->batchchanged(); //[/UserButtonCode_addButton] } else if (buttonThatWasClicked == reconstructButton) { //[UserButtonCode_reconstructButton] -- add your button handler code here.. - m_pRbmComponent->redrawReconstruction(); + m_pRbmComponentCurr->redrawReconstruction(); //[/UserButtonCode_reconstructButton] } else if (buttonThatWasClicked == ShakeButton) { //[UserButtonCode_ShakeButton] -- add your button handler code here.. - m_weights->shuffle(weightInitLabel->getText().getFloatValue()); - m_pRbmComponent->redrawWeights(); - m_pRbmComponent->redrawReconstruction(); + m_weightsCurr->shuffle(weightInitLabel->getText().getFloatValue()); + m_pRbmComponentCurr->batchchanged(); + m_pRbmComponentCurr->redrawWeights(); + m_pRbmComponentCurr->redrawReconstruction(); //[/UserButtonCode_ShakeButton] } else if (buttonThatWasClicked == testButton) { //[UserButtonCode_testButton] -- add your button handler code here.. - m_pRbmComponent->copyReconstructionToTraining(); - m_pRbmComponent->redrawReconstruction(); + m_pRbmComponentCurr->copyReconstructionToTraining(); + m_pRbmComponentCurr->redrawReconstruction(); //[/UserButtonCode_testButton] } else if (buttonThatWasClicked == createButton) @@ -495,7 +506,8 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) else if (buttonThatWasClicked == loadButton) { //[UserButtonCode_loadButton] -- add your button handler code here.. - load(); + m_rbmSelect->setSelectedId(1, sendNotification); + create(getBaseDir()); //[/UserButtonCode_loadButton] } else if (buttonThatWasClicked == saveButton) @@ -531,48 +543,48 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) else if (buttonThatWasClicked == reconstructEquButton) { //[UserButtonCode_reconstructEquButton] -- add your button handler code here.. - m_pRbmComponent->redrawReconstruction(); + m_pRbmComponentCurr->redrawReconstruction(); //[/UserButtonCode_reconstructEquButton] } else if (buttonThatWasClicked == rbmDoRaoBlackwellToggleButton) { //[UserButtonCode_rbmDoRaoBlackwellToggleButton] -- add your button handler code here.. - m_pRbmComponent->setDoRaoBlackwell(buttonThatWasClicked->getToggleState()); + m_pRbmComponentCurr->setDoRaoBlackwell(buttonThatWasClicked->getToggleState()); //[/UserButtonCode_rbmDoRaoBlackwellToggleButton] } else if (buttonThatWasClicked == rbmReduceVarianceToggleButton) { //[UserButtonCode_rbmReduceVarianceToggleButton] -- add your button handler code here.. - m_pRbmComponent->setUseProbsForHiddenReconstruction(buttonThatWasClicked->getToggleState()); + m_pRbmComponentCurr->setUseProbsForHiddenReconstruction(buttonThatWasClicked->getToggleState()); //[/UserButtonCode_rbmReduceVarianceToggleButton] } else if (buttonThatWasClicked == rbmUseVisibleGaussianToggleButton) { //[UserButtonCode_rbmUseVisibleGaussianToggleButton] -- add your button handler code here.. - m_pRbmComponent->setUseVisibleGaussian(buttonThatWasClicked->getToggleState()); - m_pRbmComponent->redrawReconstruction(); + m_pRbmComponentCurr->setUseVisibleGaussian(buttonThatWasClicked->getToggleState()); + m_pRbmComponentCurr->redrawReconstruction(); //[/UserButtonCode_rbmUseVisibleGaussianToggleButton] } else if (buttonThatWasClicked == rbmDoSparseToggleButton) { //[UserButtonCode_rbmDoSparseToggleButton] -- add your button handler code here.. - m_pRbmComponent->setDoSparse(buttonThatWasClicked->getToggleState()); + m_pRbmComponentCurr->setDoSparse(buttonThatWasClicked->getToggleState()); //[/UserButtonCode_rbmDoSparseToggleButton] } else if (buttonThatWasClicked == rbmLearnVarianceButton) { //[UserButtonCode_rbmLearnVarianceButton] -- add your button handler code here.. - m_pRbmComponent->setDoLearnVariance(buttonThatWasClicked->getToggleState()); + m_pRbmComponentCurr->setDoLearnVariance(buttonThatWasClicked->getToggleState()); if (!buttonThatWasClicked->getToggleState()) { - m_pRbmComponent->setSigma(sigmaLabel->getText().getFloatValue()); + m_pRbmComponentCurr->setSigma(sigmaLabel->getText().getFloatValue()); } //[/UserButtonCode_rbmLearnVarianceButton] } else if (buttonThatWasClicked == rbmNormalizeDataToggleButton) { //[UserButtonCode_rbmNormalizeDataToggleButton] -- add your button handler code here.. - m_pRbmComponent->setNormalizeData(buttonThatWasClicked->getToggleState()); + m_pRbmComponentCurr->setNormalizeData(buttonThatWasClicked->getToggleState()); //[/UserButtonCode_rbmNormalizeDataToggleButton] } @@ -584,23 +596,23 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved) { //[UsersliderValueChanged_Pre] //[/UsersliderValueChanged_Pre] - + if (sliderThatWasMoved == patterSlider) { //[UserSliderCode_patterSlider] -- add your slider handling code here.. - m_pRbmComponent->selectTraining((int)sliderThatWasMoved->getValue()); + m_pRbmComponentCurr->setTrainingIndex((int)sliderThatWasMoved->getValue()); //[/UserSliderCode_patterSlider] } else if (sliderThatWasMoved == WeightsSlider) { //[UserSliderCode_WeightsSlider] -- add your slider handling code here.. - m_pRbmComponent->selectWeights((int)sliderThatWasMoved->getValue()); + m_pRbmComponentCurr->setWeightsIndex((int)sliderThatWasMoved->getValue()); //[/UserSliderCode_WeightsSlider] } else if (sliderThatWasMoved == numGibbsSlider) { //[UserSliderCode_numGibbsSlider] -- add your slider handling code here.. - m_pRbmComponent->setNumGibbs((uint32_t)sliderThatWasMoved->getValue()); + m_pRbmComponentCurr->setNumGibbs((uint32_t)sliderThatWasMoved->getValue()); //[/UserSliderCode_numGibbsSlider] } else if (sliderThatWasMoved == m_progressBarSlider) @@ -626,7 +638,7 @@ void MainComponent::labelTextChanged (Label* labelThatHasChanged) else if (labelThatHasChanged == learningRateLabel) { //[UserLabelCode_learningRateLabel] -- add your label text handling code here.. - m_pRbmComponent->setMuWeights(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setMuWeights(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_learningRateLabel] } else if (labelThatHasChanged == numVisibleLabel) @@ -652,44 +664,44 @@ void MainComponent::labelTextChanged (Label* labelThatHasChanged) else if (labelThatHasChanged == lambdaLabel) { //[UserLabelCode_lambdaLabel] -- add your label text handling code here.. - m_pRbmComponent->setLambda(labelThatHasChanged->getText().getFloatValue()); - m_pRbmComponent->redrawReconstruction(); + m_pRbmComponentCurr->setLambda(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->redrawReconstruction(); //[/UserLabelCode_lambdaLabel] } else if (labelThatHasChanged == sigmaLabel) { //[UserLabelCode_sigmaLabel] -- add your label text handling code here.. - m_pRbmComponent->setSigma(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setSigma(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_sigmaLabel] } else if (labelThatHasChanged == sparsityLabel) { //[UserLabelCode_sparsityLabel] -- add your label text handling code here.. - m_pRbmComponent->setSparsity(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setSparsity(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_sparsityLabel] } else if (labelThatHasChanged == sigmaDecayLabel) { //[UserLabelCode_sigmaDecayLabel] -- add your label text handling code here.. - m_pRbmComponent->setSigmaDecay(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setSigmaDecay(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_sigmaDecayLabel] } else if (labelThatHasChanged == weightDecayLabel) { //[UserLabelCode_weightDecayLabel] -- add your label text handling code here.. - m_pRbmComponent->setWeightDecay(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setWeightDecay(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_weightDecayLabel] } else if (labelThatHasChanged == momentumLabel) { //[UserLabelCode_momentumLabel] -- add your label text handling code here.. - m_pRbmComponent->setMomentum(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setMomentum(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_momentumLabel] } else if (labelThatHasChanged == sparsityLearningRateLabel) { //[UserLabelCode_sparsityLearningRateLabel] -- add your label text handling code here.. - m_pRbmComponent->setMuSparsity(labelThatHasChanged->getText().getFloatValue()); + m_pRbmComponentCurr->setMuSparsity(labelThatHasChanged->getText().getFloatValue()); //[/UserLabelCode_sparsityLearningRateLabel] } else if (labelThatHasChanged == weightInitLabel) @@ -702,6 +714,29 @@ void MainComponent::labelTextChanged (Label* labelThatHasChanged) //[/UserlabelTextChanged_Post] } +void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged) +{ + //[UsercomboBoxChanged_Pre] + //[/UsercomboBoxChanged_Pre] + + if (comboBoxThatHasChanged == m_rbmSelect) + { + //[UserComboBoxCode_m_rbmSelect] -- add your combo box handling code here.. + m_pRbmComponentCurr = m_pRbmComponent[m_rbmSelect->getSelectedId()-1]; + m_weightsCurr = m_weights[m_rbmSelect->getSelectedId()-1]; + + if (m_pRbmComponentCurr == nullptr) + { + create(); + } + updateControls(); + //[/UserComboBoxCode_m_rbmSelect] + } + + //[UsercomboBoxChanged_Post] + //[/UsercomboBoxChanged_Post] +} + void MainComponent::mouseMove (const MouseEvent& e) { //[UserCode_mouseMove] -- Add your code here... @@ -755,61 +790,44 @@ void MainComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails //[MiscUserCode] You can add your own definitions of your custom methods or any other code here... -void MainComponent::load () -{ - create(getBaseDir()); -} - void MainComponent::save () { - m_weights->save((String(getBaseDir() + String(".weights.dat"))).toUTF8()); + m_weights[0]->save((String(getBaseDir() + String(".weights.dat"))).toUTF8()); } void MainComponent::create(juce::String const &projectName) { - m_weights = nullptr; - m_pRbmComponent = nullptr; + size_t id = m_rbmSelect->getSelectedId()-1; + if ((m_pRbmComponent[id] == nullptr) and ((id+1) < DBN_SIZE)) + m_rbmSelect->addItem(String(id+1), id+2); - if (!projectName.isEmpty()) + m_weights[id] = nullptr; + m_pRbmComponent[id] = nullptr; + + if (id == 0) { - m_weights = new Weights((String(projectName + String(".weights.dat"))).toUTF8()); - loadTraining((String(projectName + String(".trainingStates.dat"))).toUTF8()); + if (!projectName.isEmpty()) + { + m_weights[id] = new Weights((String(projectName + String(".weights.dat"))).toUTF8()); + loadTraining((String(projectName + String(".trainingStates.dat"))).toUTF8()); + } + else + { + m_weights[id] = new Weights(numVisibleLabel->getText().getIntValue(), numVisibleYLabel->getText().getIntValue(), numHiddenLabel->getText().getIntValue()); + } + addAndMakeVisible(m_pRbmComponent[id] = new RbmComponent(*m_weights[id], m_layers.data(), *this)); } else { - m_weights = new Weights(numVisibleLabel->getText().getIntValue(), numVisibleYLabel->getText().getIntValue(), numHiddenLabel->getText().getIntValue()); + m_weights[id] = new Weights(m_weights[id-1]->getNumHidden(), 1, numHiddenLabel->getText().getIntValue()); + addAndMakeVisible(m_pRbmComponent[id] = new RbmComponent(*m_weights[id], m_pRbmComponent[id-1]->getHiddenBatch(), *this)); } - size_t vNumX = m_weights->getNumVisibleX(); - size_t vNumY = m_weights->getNumVisibleY(); - size_t hNum = m_weights->getNumHidden(); - WeightsSlider->setRange(0, hNum-1, 1); - - numVisibleLabel->setText(String(vNumX), dontSendNotification ); - numVisibleYLabel->setText(String(vNumY), dontSendNotification ); - numHiddenLabel->setText(String(hNum), dontSendNotification ); - - addAndMakeVisible(m_pRbmComponent = new RbmComponent(*m_weights, m_layers.data(), *this)); - m_pRbmComponent->setBounds (16, 16, 430, 130); - - m_pRbmComponent->setDoRaoBlackwell(rbmDoRaoBlackwellToggleButton->getToggleState()); - m_pRbmComponent->setUseProbsForHiddenReconstruction(rbmReduceVarianceToggleButton->getToggleState()); - m_pRbmComponent->setUseVisibleGaussian(rbmUseVisibleGaussianToggleButton->getToggleState()); - m_pRbmComponent->setDoSparse(rbmDoSparseToggleButton->getToggleState()); - - m_pRbmComponent->setLambda(lambdaLabel->getText().getFloatValue()); - m_pRbmComponent->setSigmaDecay(sigmaDecayLabel->getText().getFloatValue()); - m_pRbmComponent->setWeightDecay(weightDecayLabel->getText().getFloatValue()); - m_pRbmComponent->setSparsity(sparsityLabel->getText().getFloatValue()); - m_pRbmComponent->setNumGibbs((uint32_t)numGibbsSlider->getValue()); - m_pRbmComponent->setDoLearnVariance(rbmLearnVarianceButton->getToggleState()); - if (!rbmLearnVarianceButton->getToggleState()) - { - m_pRbmComponent->setSigma(sigmaLabel->getText().getFloatValue()); - } - - m_pRbmComponent->redrawReconstruction(); - m_pRbmComponent->selectWeights((int)WeightsSlider->getValue()); - m_pRbmComponent->selectTraining((int)patterSlider->getValue()); + m_pRbmComponentCurr = m_pRbmComponent[id]; + m_weightsCurr = m_weights[id]; + m_pRbmComponentCurr->setBounds (16, 140*id+16, 430, 130); + m_pRbmComponentCurr->batchchanged(); + m_pRbmComponentCurr->redrawReconstruction(); + updateControls(); } void MainComponent::destroy() @@ -827,7 +845,7 @@ const juce::String& MainComponent::getBaseDir() void MainComponent::run() { // trainButton->setEnabled(false); - m_pRbmComponent->train(numEpochslabel->getText().getIntValue()); + m_pRbmComponentCurr->train(numEpochslabel->getText().getIntValue()); // trainButton->setEnabled(true); } @@ -841,6 +859,34 @@ void MainComponent::onRbmEpochTrained(size_t progressPercent) m_progressBarSlider->setValue(progressPercent); } +void MainComponent::updateControls() +{ + + rbmDoRaoBlackwellToggleButton->setToggleState(m_pRbmComponentCurr->params().m_doRaoBlackwell, dontSendNotification); + rbmReduceVarianceToggleButton->setToggleState(m_pRbmComponentCurr->params().m_useProbsForHiddenReconstruction, dontSendNotification); + rbmUseVisibleGaussianToggleButton->setToggleState(m_pRbmComponentCurr->params().m_useVisibleGaussian, dontSendNotification); + rbmDoSparseToggleButton->setToggleState(m_pRbmComponentCurr->params().m_doSparse, dontSendNotification); + rbmLearnVarianceButton->setToggleState(m_pRbmComponentCurr->params().m_doLearnVariance, dontSendNotification); + + lambdaLabel->setText(String(m_pRbmComponentCurr->params().m_lambda), dontSendNotification); + sigmaDecayLabel->setText(String(m_pRbmComponentCurr->params().m_sigmaDecay), dontSendNotification); + weightDecayLabel->setText(String(m_pRbmComponentCurr->params().m_weightDecay), dontSendNotification); + sparsityLabel->setText(String(m_pRbmComponentCurr->params().m_sparsity), dontSendNotification); + sigmaLabel->setText(String(m_pRbmComponentCurr->params().m_constantSigma), dontSendNotification); + sparsityLearningRateLabel->setText(String(m_pRbmComponentCurr->params().m_muSparsity), dontSendNotification); + learningRateLabel->setText(String(m_pRbmComponentCurr->params().m_muWeights), dontSendNotification); + momentumLabel->setText(String(m_pRbmComponentCurr->params().m_momentum), dontSendNotification); + numVisibleLabel->setText(String(m_weightsCurr->getNumVisibleX()), dontSendNotification ); + numVisibleYLabel->setText(String(m_weightsCurr->getNumVisibleY()), dontSendNotification ); + numHiddenLabel->setText(String(m_weightsCurr->getNumHidden()), dontSendNotification ); + + WeightsSlider->setRange(0, m_weightsCurr->getNumHidden()-1, 1); + numGibbsSlider->setValue(m_pRbmComponentCurr->params().m_numGibbs); + WeightsSlider->setValue(m_pRbmComponentCurr->getWeightsIndex()); + patterSlider->setValue(m_pRbmComponentCurr->getTrainingIndex()); + +} + //[/MiscUserCode] @@ -854,8 +900,8 @@ void MainComponent::onRbmEpochTrained(size_t progressPercent) BEGIN_JUCER_METADATA @@ -1036,6 +1082,9 @@ BEGIN_JUCER_METADATA + END_JUCER_METADATA diff --git a/Source/MainComponent.h b/Source/MainComponent.h index b7cbe56..f3c6aa5 100644 --- a/Source/MainComponent.h +++ b/Source/MainComponent.h @@ -43,7 +43,8 @@ class MainComponent : public Component, public Thread, public ButtonListener, public SliderListener, - public LabelListener + public LabelListener, + public ComboBoxListener { public: //============================================================================== @@ -61,6 +62,7 @@ public: void buttonClicked (Button* buttonThatWasClicked); void sliderValueChanged (Slider* sliderThatWasMoved); void labelTextChanged (Label* labelThatHasChanged); + void comboBoxChanged (ComboBox* comboBoxThatHasChanged); void mouseMove (const MouseEvent& e); void mouseEnter (const MouseEvent& e); void mouseExit (const MouseEvent& e); @@ -74,10 +76,12 @@ public: private: //[UserVariables] -- You can add your own custom variables in this section. - ScopedPointer m_weights; - ScopedPointer m_pRbmComponent; + static const size_t DBN_SIZE = 4; + ScopedPointer m_weights[DBN_SIZE]; + ScopedPointer m_pRbmComponent[DBN_SIZE]; + Weights *m_weightsCurr; + RbmComponent *m_pRbmComponentCurr; LayerArray m_layers; - void load(); void save(); void create(juce::String const &projectName=""); void destroy(); @@ -88,22 +92,25 @@ private: { m_layers.clear(); } - + void loadTraining(const char *pFilename) { m_layers.clear(); m_layers.load(pFilename); } - + void saveTraining(const char *pFilename) { m_layers.save(pFilename); } - + void removeTrainingAt(size_t index) { m_layers.removeAt(index); } + + void updateControls(); + //[/UserVariables] //============================================================================== @@ -144,6 +151,7 @@ private: ScopedPointer