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