diff --git a/Makefile b/Makefile
index e097081..9f8f349 100644
--- a/Makefile
+++ b/Makefile
@@ -2,7 +2,7 @@ CONFIG ?= release
BUILD_DIR := ./build/${CONFIG}
JUCE_ROOT := ../JUCE-3.1.1
-GPP_SRCS := Main.cpp Rbm.cpp Layer.cpp Stack.cpp MainComponent.cpp
+GPP_SRCS := gui.cpp Rbm.cpp Layer.cpp Stack.cpp MainComponent.cpp RbmComponent.cpp DrawComponent.cpp
GPP_OBJS := $(addprefix ${BUILD_DIR}/, $(subst .cpp,.o,$(GPP_SRCS)))
GCC_SRCS := noise.c
GCC_OBJS := $(addprefix ${BUILD_DIR}/, $(subst .c,.o,$(GCC_SRCS)))
diff --git a/source/DrawComponent.cpp b/source/DrawComponent.cpp
new file mode 100644
index 0000000..be94cb6
--- /dev/null
+++ b/source/DrawComponent.cpp
@@ -0,0 +1,292 @@
+/*
+ ==============================================================================
+
+ This is an automatically generated GUI class created by the Introjucer!
+
+ Be careful when adding custom code to these files, as only the code within
+ the "//[xyz]" and "//[/xyz]" sections will be retained when the file is loaded
+ and re-saved.
+
+ Created with Introjucer version: 3.1.0
+
+ ------------------------------------------------------------------------------
+
+ The Introjucer is part of the JUCE library - "Jules' Utility Class Extensions"
+ Copyright 2004-13 by Raw Material Software Ltd.
+
+ ==============================================================================
+*/
+
+//[Headers] You can add your own extra header files here...
+#include "noise.h"
+void mylog(const char* format, ...);
+//[/Headers]
+
+#include "DrawComponent.hpp"
+
+
+//[MiscUserDefs] You can add your own user definitions and misc code here...
+//[/MiscUserDefs]
+
+//==============================================================================
+DrawComponent::DrawComponent (int width, int height, float offset, float scale)
+ : m_pListener(nullptr)
+{
+
+ //[UserPreSize]
+ m_width = width;
+ m_height = height;
+ m_offset = offset;
+ m_scale = scale;
+ m_scaleX = 1.0;
+ m_scaleY = 1.0;
+ m_pG = nullptr;
+
+ //[/UserPreSize]
+
+ setSize (80, 80);
+
+
+ //[Constructor] You can add your own custom stuff here..
+// m_pG = new Graphics(m_image);
+ m_data.resize(width*height);
+
+ m_pG->setImageResamplingQuality(Graphics::lowResamplingQuality);
+ clear();
+ //[/Constructor]
+}
+
+DrawComponent::~DrawComponent()
+{
+ //[Destructor_pre]. You can add your own custom destruction code here..
+ //[/Destructor_pre]
+
+
+
+ //[Destructor]. You can add your own custom destruction code here..
+ m_pG = nullptr;
+ m_pListener = nullptr;
+ m_data.resize(0);
+ //[/Destructor]
+}
+
+//==============================================================================
+void DrawComponent::paint (Graphics& g)
+{
+ //[UserPrePaint] Add your own custom painting code here..
+ AffineTransform t;
+
+ g.drawImage(m_image, 0, 0, getWidth(), getHeight(), 0, 0, getWidth(), getHeight(), false);
+// g.drawImageTransformed(*m_pImage, t.rotated(8), false);
+ return;
+ //[/UserPrePaint]
+
+ g.fillAll (Colours::grey);
+
+ //[UserPaint] Add your own custom painting code here..
+ //[/UserPaint]
+}
+
+void DrawComponent::resized()
+{
+ //[UserResized] Add your own custom resize handling here..
+// mylog("%s (%d,%d)\n", __func__, getWidth(), getHeight());
+ m_scaleX = (float)getWidth()/m_width;
+ m_scaleY = (float)getHeight()/m_height;
+ m_image = Image(Image::RGB , getWidth(), getHeight(), true);
+ m_pG = nullptr;
+ m_pG = new Graphics(m_image);
+
+ repaint();
+ //[/UserResized]
+}
+
+void DrawComponent::mouseMove (const MouseEvent& e)
+{
+ //[UserCode_mouseMove] -- Add your code here...
+// mylogMessage(String("MouseMove"));
+ //[/UserCode_mouseMove]
+}
+
+void DrawComponent::mouseEnter (const MouseEvent& e)
+{
+ //[UserCode_mouseEnter] -- Add your code here...
+// mylog("%s\n", __func__);
+ //[/UserCode_mouseEnter]
+}
+
+void DrawComponent::mouseExit (const MouseEvent& e)
+{
+ //[UserCode_mouseExit] -- Add your code here...
+// mylog("%s\n", __func__);
+ //[/UserCode_mouseExit]
+}
+
+void DrawComponent::mouseDown (const MouseEvent& e)
+{
+ //[UserCode_mouseDown] -- Add your code here...
+// mylog("%s x:%d, y:%d\n", __func__, e.x, e.y);
+ if (e.mods.isRightButtonDown())
+ {
+ clear();
+ }
+ else
+ {
+ drawAt(e.x, e.y, true);
+ }
+ if (m_pListener)
+ m_pListener->onDraw(*this);
+ //[/UserCode_mouseDown]
+}
+
+void DrawComponent::mouseDrag (const MouseEvent& e)
+{
+ //[UserCode_mouseDrag] -- Add your code here...
+// mylog("%s x:%d, y:%d\n", __func__, e.x, e.y);
+ if (e.mods.isLeftButtonDown())
+ {
+ drawAt(e.x, e.y, false);
+ if (m_pListener)
+ m_pListener->onDraw(*this);
+ }
+ //[/UserCode_mouseDrag]
+}
+
+void DrawComponent::mouseUp (const MouseEvent& e)
+{
+ //[UserCode_mouseUp] -- Add your code here...
+// mylog("%s\n", __func__);
+ //[/UserCode_mouseUp]
+}
+
+void DrawComponent::mouseDoubleClick (const MouseEvent& e)
+{
+ //[UserCode_mouseDoubleClick] -- Add your code here...
+// mylog("%s\n", __func__);
+ //[/UserCode_mouseDoubleClick]
+}
+
+void DrawComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel)
+{
+ //[UserCode_mouseWheelMove] -- Add your code here...
+// mylog("%s\n", __func__);
+ //[/UserCode_mouseWheelMove]
+}
+
+
+
+//[MiscUserCode] You can add your own definitions of your custom methods or any other code here...
+void DrawComponent::setListener(DrawListener *pListener)
+{
+ m_pListener = pListener;
+}
+
+void DrawComponent::drawAt(int x, int y, bool setColor)
+{
+ int index;
+ double fx = (double)x / m_scaleX;
+ double fy = (double)y / m_scaleY;
+
+ index = (int)(x/m_scaleX) + m_width*(int)(y/m_scaleY);
+// mylog("Write(%d)\n", index);
+
+ if (index >= 0)
+ {
+ if (index < m_width*m_height)
+ {
+ if (setColor)
+ {
+ if (m_data[index] > 0.0)
+ {
+ m_currData = 0.0;
+ m_currColor = (Colours::black);
+ }
+ else
+ {
+ m_currData = 1.0;
+ m_currColor = (Colours::white);
+ }
+ }
+ m_data[index] = m_currData;
+ m_pG->setColour (m_currColor);
+ }
+ }
+
+ x = (int)fx * m_scaleX;
+ y = (int)fy * m_scaleY;
+ m_pG->fillRect((float)x, (float)y, m_scaleX, m_scaleY);
+ repaint();
+}
+
+void DrawComponent::clear()
+{
+ m_data.resize(m_width*m_height);
+ m_data = arma::zeros(m_width*m_height);
+ DrawData();
+}
+
+arma::mat& DrawComponent::getData ()
+{
+ return m_data;
+}
+
+void DrawComponent::DrawData ()
+{
+
+ double max = abs(m_data).max();
+
+ if (max < 1.0)
+ {
+ max = 1.0;
+ }
+ double scale = m_scale/max;
+
+ for (int i=0; i < m_height; i++)
+ {
+ for (int j=0; j < m_width; j++)
+ {
+ m_pG->setColour(Colour(Colours::white).greyLevel(m_offset + scale*m_data[i*m_width + j]));
+ m_pG->fillRect(m_scaleX*j, m_scaleY*i, m_scaleX, m_scaleY);
+ }
+ }
+ repaint();
+
+}
+
+//[/MiscUserCode]
+
+
+//==============================================================================
+#if 0
+/* -- Introjucer information section --
+
+ This is where the Introjucer stores the metadata that describe this GUI layout, so
+ make changes in here at your peril!
+
+BEGIN_JUCER_METADATA
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+END_JUCER_METADATA
+*/
+#endif
+
+
+//[EndFile] You can add extra defines here...
+//[/EndFile]
diff --git a/source/DrawComponent.hpp b/source/DrawComponent.hpp
new file mode 100644
index 0000000..978c83c
--- /dev/null
+++ b/source/DrawComponent.hpp
@@ -0,0 +1,102 @@
+/*
+ ==============================================================================
+
+ This is an automatically generated GUI class created by the Introjucer!
+
+ Be careful when adding custom code to these files, as only the code within
+ the "//[xyz]" and "//[/xyz]" sections will be retained when the file is loaded
+ and re-saved.
+
+ Created with Introjucer version: 3.1.0
+
+ ------------------------------------------------------------------------------
+
+ The Introjucer is part of the JUCE library - "Jules' Utility Class Extensions"
+ Copyright 2004-13 by Raw Material Software Ltd.
+
+ ==============================================================================
+*/
+
+#ifndef __JUCE_HEADER_7594B61AA9F7A49E__
+#define __JUCE_HEADER_7594B61AA9F7A49E__
+
+//[Headers] -- You can add your own extra header files here --
+#include
+#include "JuceHeader.h"
+class DrawComponent;
+class DrawListener
+{
+public:
+ DrawListener() {}
+ virtual ~DrawListener() {}
+ virtual void onDraw(DrawComponent &obj) = 0;
+};
+
+//[/Headers]
+
+
+
+//==============================================================================
+/**
+ //[Comments]
+ An auto-generated component, created by the Introjucer.
+
+ Describe your class and how it works here!
+ //[/Comments]
+*/
+class DrawComponent : public Component
+{
+public:
+ //==============================================================================
+ DrawComponent (int width, int height, float offset=0.0, float scale=1.0);
+ ~DrawComponent();
+
+ //==============================================================================
+ //[UserMethods] -- You can add your own custom methods in this section.
+ void setListener(DrawListener *pListener);
+ void drawAt(int x, int y, bool setColor);
+ arma::mat& getData();
+ void clear();
+ void DrawData();
+ //[/UserMethods]
+
+ void paint (Graphics& g);
+ void resized();
+ void mouseMove (const MouseEvent& e);
+ void mouseEnter (const MouseEvent& e);
+ void mouseExit (const MouseEvent& e);
+ void mouseDown (const MouseEvent& e);
+ void mouseDrag (const MouseEvent& e);
+ void mouseUp (const MouseEvent& e);
+ void mouseDoubleClick (const MouseEvent& e);
+ void mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel);
+
+
+
+private:
+ //[UserVariables] -- You can add your own custom variables in this section.
+ DrawListener *m_pListener;
+ int m_width;
+ int m_height;
+ float m_offset;
+ float m_scale;
+ float m_scaleX;
+ float m_scaleY;
+ ScopedPointerm_pG;
+ arma::mat m_data;
+ Image m_image;
+ double m_currData;
+ Colour m_currColor;
+ //[/UserVariables]
+
+ //==============================================================================
+
+
+ //==============================================================================
+ JUCE_DECLARE_NON_COPYABLE_WITH_LEAK_DETECTOR (DrawComponent)
+};
+
+//[EndFile] You can add extra defines here...
+//[/EndFile]
+
+#endif // __JUCE_HEADER_7594B61AA9F7A49E__
diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp
new file mode 100644
index 0000000..929a6ec
--- /dev/null
+++ b/source/MainComponent.cpp
@@ -0,0 +1,910 @@
+/*
+ ==============================================================================
+
+ This is an automatically generated GUI class created by the Introjucer!
+
+ Be careful when adding custom code to these files, as only the code within
+ the "//[xyz]" and "//[/xyz]" sections will be retained when the file is loaded
+ and re-saved.
+
+ Created with Introjucer version: 3.1.0
+
+ ------------------------------------------------------------------------------
+
+ The Introjucer is part of the JUCE library - "Jules' Utility Class Extensions"
+ Copyright 2004-13 by Raw Material Software Ltd.
+
+ ==============================================================================
+*/
+
+//[Headers] You can add your own extra header files here...
+//[/Headers]
+
+#include
+#include
+#include
+#include "MainComponent.hpp"
+
+
+//[MiscUserDefs] You can add your own user definitions and misc code here...
+void mylog(const char* format, ...)
+{
+ char log_buf[257];
+ va_list argptr;
+ va_start(argptr, format);
+ vprintf(format, argptr);
+ va_end(argptr);
+}
+
+//[/MiscUserDefs]
+
+//==============================================================================
+MainComponent::MainComponent ()
+: Thread("RBM")
+, m_stack(nullptr)
+, m_pLayer(nullptr)
+, m_trainingIndex(0)
+, m_weightIndex(0)
+, m_toolTipWindow(this)
+{
+ addAndMakeVisible (trainButton = new TextButton ("Train button"));
+ trainButton->setButtonText (TRANS("Train"));
+ trainButton->addListener (this);
+
+ addAndMakeVisible (addButton = new TextButton ("Add button"));
+ addButton->setButtonText (TRANS("Add"));
+ addButton->addListener (this);
+
+ addAndMakeVisible (patterSlider = new Slider ("Pattern slider"));
+ patterSlider->setTooltip (TRANS("Training index"));
+ patterSlider->setRange (0, 0, 1);
+ patterSlider->setSliderStyle (Slider::LinearBar);
+ patterSlider->setTextBoxStyle (Slider::TextBoxLeft, false, 80, 20);
+ patterSlider->addListener (this);
+
+ addAndMakeVisible (ShakeButton = new TextButton ("Shake button"));
+ ShakeButton->setButtonText (TRANS("Shake"));
+ ShakeButton->addListener (this);
+
+ addAndMakeVisible (WeightsSlider = new Slider ("Weights slider"));
+ WeightsSlider->setTooltip (TRANS("Weight index"));
+ WeightsSlider->setRange (0, 0, 1);
+ WeightsSlider->setSliderStyle (Slider::LinearBar);
+ WeightsSlider->setTextBoxStyle (Slider::TextBoxLeft, false, 80, 20);
+ WeightsSlider->addListener (this);
+
+ addAndMakeVisible (numEpochslabel = new Label ("Num Epochs label",
+ TRANS("1000")));
+ numEpochslabel->setFont (Font (15.00f, Font::plain));
+ numEpochslabel->setJustificationType (Justification::centred);
+ numEpochslabel->setEditable (true, true, false);
+ numEpochslabel->setColour (TextEditor::textColourId, Colours::black);
+ numEpochslabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ numEpochslabel->addListener (this);
+
+ addAndMakeVisible (learningRateLabel = new Label ("Learning Rate label",
+ TRANS("0.01")));
+ learningRateLabel->setFont (Font (15.00f, Font::plain));
+ learningRateLabel->setJustificationType (Justification::centred);
+ learningRateLabel->setEditable (true, true, false);
+ learningRateLabel->setColour (TextEditor::textColourId, Colours::black);
+ learningRateLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ learningRateLabel->addListener (this);
+
+ addAndMakeVisible (numVisibleLabel = new Label ("Num Visible label",
+ TRANS("16")));
+ numVisibleLabel->setTooltip (TRANS("Number of visible units X"));
+ numVisibleLabel->setFont (Font (15.00f, Font::plain));
+ numVisibleLabel->setJustificationType (Justification::centred);
+ numVisibleLabel->setEditable (true, true, false);
+ numVisibleLabel->setColour (TextEditor::textColourId, Colours::black);
+ numVisibleLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ numVisibleLabel->addListener (this);
+
+ addAndMakeVisible (numHiddenLabel = new Label ("Num Hidden label",
+ TRANS("64")));
+ numHiddenLabel->setTooltip (TRANS("Number of hidden units"));
+ numHiddenLabel->setFont (Font (15.00f, Font::plain));
+ numHiddenLabel->setJustificationType (Justification::centred);
+ numHiddenLabel->setEditable (true, true, false);
+ numHiddenLabel->setColour (TextEditor::textColourId, Colours::black);
+ numHiddenLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ numHiddenLabel->addListener (this);
+
+ addAndMakeVisible (createButton = new TextButton ("Create button"));
+ createButton->setButtonText (TRANS("Create"));
+ createButton->addListener (this);
+
+ addAndMakeVisible (projectNameLabel = new Label ("Project Name label",
+ TRANS("test")));
+ projectNameLabel->setFont (Font (15.00f, Font::plain));
+ projectNameLabel->setJustificationType (Justification::centred);
+ projectNameLabel->setEditable (true, true, false);
+ projectNameLabel->setColour (TextEditor::textColourId, Colours::black);
+ projectNameLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ projectNameLabel->addListener (this);
+
+ addAndMakeVisible (loadButton = new TextButton ("Load button"));
+ loadButton->setButtonText (TRANS("Load"));
+ loadButton->addListener (this);
+
+ addAndMakeVisible (saveButton = new TextButton ("Save button"));
+ saveButton->setButtonText (TRANS("Save"));
+ saveButton->addListener (this);
+
+ addAndMakeVisible (numVisibleYLabel = new Label ("Num Visible Y label",
+ TRANS("16")));
+ numVisibleYLabel->setTooltip (TRANS("Number of visible units Y"));
+ numVisibleYLabel->setFont (Font (15.00f, Font::plain));
+ numVisibleYLabel->setJustificationType (Justification::centred);
+ numVisibleYLabel->setEditable (true, true, false);
+ numVisibleYLabel->setColour (TextEditor::textColourId, Colours::black);
+ numVisibleYLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ numVisibleYLabel->addListener (this);
+
+ addAndMakeVisible (loadTrainingButton = new TextButton ("Load Training button"));
+ loadTrainingButton->setButtonText (TRANS("Load T"));
+ loadTrainingButton->addListener (this);
+
+ addAndMakeVisible (saveTrainingButton = new TextButton ("Save Training button"));
+ saveTrainingButton->setButtonText (TRANS("Save T"));
+ saveTrainingButton->addListener (this);
+
+ addAndMakeVisible (clearTrainingButton = new TextButton ("Clear Training button"));
+ clearTrainingButton->setButtonText (TRANS("Clear T"));
+ clearTrainingButton->addListener (this);
+
+ addAndMakeVisible (removeTrainingButton = new TextButton ("Remove Training button"));
+ removeTrainingButton->setButtonText (TRANS("Remove T"));
+ removeTrainingButton->addListener (this);
+
+ addAndMakeVisible (numGibbsSlider = new Slider ("Num. Gibbs slider"));
+ numGibbsSlider->setTooltip (TRANS("Number of gibbs samples"));
+ numGibbsSlider->setRange (1, 101, 1);
+ numGibbsSlider->setSliderStyle (Slider::LinearBar);
+ numGibbsSlider->setTextBoxStyle (Slider::TextBoxLeft, false, 80, 20);
+ numGibbsSlider->addListener (this);
+
+ addAndMakeVisible (rbmDoRaoBlackwellToggleButton = new ToggleButton ("rbmDoRaoBlackwell toggle button"));
+ rbmDoRaoBlackwellToggleButton->setTooltip (TRANS("Don\'t sample hidden during learning"));
+ rbmDoRaoBlackwellToggleButton->setButtonText (TRANS("Rao-Blackwell"));
+ rbmDoRaoBlackwellToggleButton->addListener (this);
+
+ addAndMakeVisible (rbmDoSampleVisibleToggleButton = new ToggleButton ("rbmDoSampleVisible toggle button"));
+ rbmDoSampleVisibleToggleButton->setTooltip (TRANS("Sample visible during learning\n"));
+ rbmDoSampleVisibleToggleButton->setButtonText (TRANS("Sample Visible"));
+ rbmDoSampleVisibleToggleButton->addListener (this);
+
+ addAndMakeVisible (lambdaLabel = new Label ("Lambda label",
+ TRANS("1.0")));
+ lambdaLabel->setFont (Font (15.00f, Font::plain));
+ lambdaLabel->setJustificationType (Justification::centred);
+ lambdaLabel->setEditable (true, true, false);
+ lambdaLabel->setColour (TextEditor::textColourId, Colours::black);
+ lambdaLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ lambdaLabel->addListener (this);
+
+ addAndMakeVisible (sigmaLabel = new Label ("Sigma label",
+ TRANS("1.0")));
+ sigmaLabel->setFont (Font (15.00f, Font::plain));
+ sigmaLabel->setJustificationType (Justification::centred);
+ sigmaLabel->setEditable (true, true, false);
+ sigmaLabel->setColour (TextEditor::textColourId, Colours::black);
+ sigmaLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ sigmaLabel->addListener (this);
+
+ addAndMakeVisible (rbmUseVisibleGaussianToggleButton = new ToggleButton ("rbmUseVisibleGaussian toggle button"));
+ rbmUseVisibleGaussianToggleButton->setButtonText (TRANS("Use gaussian visible"));
+ rbmUseVisibleGaussianToggleButton->addListener (this);
+
+ addAndMakeVisible (rbmDoSparseToggleButton = new ToggleButton ("rbmDoSparse toggle button"));
+ rbmDoSparseToggleButton->setButtonText (TRANS("Do sparse"));
+ rbmDoSparseToggleButton->addListener (this);
+
+ addAndMakeVisible (sparsityLabel = new Label ("Sparsity label",
+ TRANS("0.05")));
+ sparsityLabel->setFont (Font (15.00f, Font::plain));
+ sparsityLabel->setJustificationType (Justification::centred);
+ sparsityLabel->setEditable (true, true, false);
+ sparsityLabel->setColour (TextEditor::textColourId, Colours::black);
+ sparsityLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ sparsityLabel->addListener (this);
+
+ addAndMakeVisible (sigmaDecayLabel = new Label ("SigmaDecay label",
+ TRANS("1.0")));
+ sigmaDecayLabel->setFont (Font (15.00f, Font::plain));
+ sigmaDecayLabel->setJustificationType (Justification::centred);
+ sigmaDecayLabel->setEditable (true, true, false);
+ sigmaDecayLabel->setColour (TextEditor::textColourId, Colours::black);
+ sigmaDecayLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ sigmaDecayLabel->addListener (this);
+
+ addAndMakeVisible (weightDecayLabel = new Label ("WeightDecay label",
+ TRANS("0")));
+ weightDecayLabel->setFont (Font (15.00f, Font::plain));
+ weightDecayLabel->setJustificationType (Justification::centred);
+ weightDecayLabel->setEditable (true, true, false);
+ weightDecayLabel->setColour (TextEditor::textColourId, Colours::black);
+ weightDecayLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ weightDecayLabel->addListener (this);
+
+ addAndMakeVisible (m_progressBarSlider = new Slider ("progressBar slider"));
+ m_progressBarSlider->setTooltip (TRANS("Progress"));
+ m_progressBarSlider->setRange (0, 100, 1);
+ m_progressBarSlider->setSliderStyle (Slider::LinearBar);
+ m_progressBarSlider->setTextBoxStyle (Slider::TextBoxLeft, true, 80, 20);
+ m_progressBarSlider->addListener (this);
+
+ addAndMakeVisible (momentumLabel = new Label ("Momentum label",
+ TRANS("0.5")));
+ momentumLabel->setFont (Font (15.00f, Font::plain));
+ momentumLabel->setJustificationType (Justification::centred);
+ momentumLabel->setEditable (true, true, false);
+ momentumLabel->setColour (TextEditor::textColourId, Colours::black);
+ momentumLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ momentumLabel->addListener (this);
+
+ addAndMakeVisible (sparsityLearningRateLabel = new Label ("SparsityLearn label",
+ TRANS("0.01")));
+ sparsityLearningRateLabel->setFont (Font (15.00f, Font::plain));
+ sparsityLearningRateLabel->setJustificationType (Justification::centred);
+ sparsityLearningRateLabel->setEditable (true, true, false);
+ sparsityLearningRateLabel->setColour (TextEditor::textColourId, Colours::black);
+ sparsityLearningRateLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ sparsityLearningRateLabel->addListener (this);
+
+ addAndMakeVisible (weightInitLabel = new Label ("WeightInit label",
+ TRANS("0.01")));
+ weightInitLabel->setFont (Font (15.00f, Font::plain));
+ weightInitLabel->setJustificationType (Justification::centred);
+ weightInitLabel->setEditable (true, true, false);
+ weightInitLabel->setColour (TextEditor::textColourId, Colours::black);
+ weightInitLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ weightInitLabel->addListener (this);
+
+ addAndMakeVisible (rbmLearnVarianceButton = new ToggleButton ("rbmLearnVariance button"));
+ rbmLearnVarianceButton->setButtonText (TRANS("Learn Variance"));
+ rbmLearnVarianceButton->addListener (this);
+
+ addAndMakeVisible (rbmNormalizeDataToggleButton = new ToggleButton ("rbmNormalizeData toggle button"));
+ 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);
+
+ addAndMakeVisible (rbmDoSampleBatch = new ToggleButton ("rbmDoSampleBatch button"));
+ rbmDoSampleBatch->setTooltip (TRANS("Sample training data during learning\n"));
+ rbmDoSampleBatch->setButtonText (TRANS("Sample training"));
+ rbmDoSampleBatch->addListener (this);
+
+ addAndMakeVisible (sizeMiniBatch = new Label ("sizeMiniBatch",
+ TRANS("100")));
+ sizeMiniBatch->setTooltip (TRANS("Mini batch size"));
+ sizeMiniBatch->setFont (Font (15.00f, Font::plain));
+ sizeMiniBatch->setJustificationType (Justification::centred);
+ sizeMiniBatch->setEditable (true, true, false);
+ sizeMiniBatch->setColour (TextEditor::textColourId, Colours::black);
+ sizeMiniBatch->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
+ sizeMiniBatch->addListener (this);
+
+ addAndMakeVisible (rbmUseHiddenGaussianToggleButton = new ToggleButton ("rbmUseHiddenGaussian toggle button"));
+ rbmUseHiddenGaussianToggleButton->setButtonText (TRANS("Use gaussian hidden"));
+ rbmUseHiddenGaussianToggleButton->addListener (this);
+
+
+ //[UserPreSize]
+ m_rbmSelect->addItem(String(0), 1);
+ m_rbmSelect->setSelectedId(1, dontSendNotification);
+ m_progressBarSlider->setValue(100);
+ //[/UserPreSize]
+
+ setSize (1200, 600);
+
+
+ //[Constructor] You can add your own custom stuff here..
+ //[/Constructor]
+}
+
+MainComponent::~MainComponent()
+{
+ //[Destructor_pre]. You can add your own custom destruction code here..
+ //[/Destructor_pre]
+
+ trainButton = nullptr;
+ addButton = nullptr;
+ patterSlider = nullptr;
+ ShakeButton = nullptr;
+ WeightsSlider = nullptr;
+ numEpochslabel = nullptr;
+ learningRateLabel = nullptr;
+ numVisibleLabel = nullptr;
+ numHiddenLabel = nullptr;
+ createButton = nullptr;
+ projectNameLabel = nullptr;
+ loadButton = nullptr;
+ saveButton = nullptr;
+ numVisibleYLabel = nullptr;
+ loadTrainingButton = nullptr;
+ saveTrainingButton = nullptr;
+ clearTrainingButton = nullptr;
+ removeTrainingButton = nullptr;
+ numGibbsSlider = nullptr;
+ rbmDoRaoBlackwellToggleButton = nullptr;
+ rbmDoSampleVisibleToggleButton = nullptr;
+ lambdaLabel = nullptr;
+ sigmaLabel = nullptr;
+ rbmUseVisibleGaussianToggleButton = nullptr;
+ rbmDoSparseToggleButton = nullptr;
+ sparsityLabel = nullptr;
+ sigmaDecayLabel = nullptr;
+ weightDecayLabel = nullptr;
+ m_progressBarSlider = nullptr;
+ momentumLabel = nullptr;
+ sparsityLearningRateLabel = nullptr;
+ weightInitLabel = nullptr;
+ rbmLearnVarianceButton = nullptr;
+ rbmNormalizeDataToggleButton = nullptr;
+ m_rbmSelect = nullptr;
+ rbmDoSampleBatch = nullptr;
+ sizeMiniBatch = nullptr;
+ rbmUseHiddenGaussianToggleButton = nullptr;
+
+
+ //[Destructor]. You can add your own custom destruction code here..
+
+ //[/Destructor]
+}
+
+//==============================================================================
+void MainComponent::paint (Graphics& g)
+{
+ //[UserPrePaint] Add your own custom painting code here..
+ //[/UserPrePaint]
+
+ g.fillAll (Colours::white);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Sigma"),
+ 560, 138, 80, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Lambda"),
+ 648, 138, 80, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Num. epochs"),
+ 551, 186, 91, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Alpha"),
+ 648, 186, 80, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Sparsity"),
+ 736, 138, 80, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Sigma decay"),
+ 824, 138, 96, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Weight decay"),
+ 824, 186, 96, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Momentum"),
+ 736, 186, 80, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Sparsity Alpha"),
+ 936, 138, 96, 14,
+ Justification::centred, true);
+
+ g.setColour (Colours::black);
+ g.setFont (Font (15.00f, Font::plain));
+ g.drawText (TRANS("Weight init"),
+ 936, 186, 96, 14,
+ Justification::centred, true);
+
+ //[UserPaint] Add your own custom painting code here..
+ //[/UserPaint]
+}
+
+void MainComponent::resized()
+{
+ trainButton->setBounds (648, 20, 72, 24);
+ addButton->setBounds (556, 20, 72, 24);
+ patterSlider->setBounds (896, 280, 220, 24);
+ ShakeButton->setBounds (740, 20, 72, 24);
+ WeightsSlider->setBounds (896, 320, 220, 24);
+ numEpochslabel->setBounds (564, 200, 72, 24);
+ learningRateLabel->setBounds (652, 200, 72, 24);
+ numVisibleLabel->setBounds (656, 248, 48, 24);
+ numHiddenLabel->setBounds (780, 248, 48, 24);
+ createButton->setBounds (728, 280, 72, 24);
+ projectNameLabel->setBounds (552, 248, 96, 24);
+ loadButton->setBounds (564, 280, 72, 24);
+ saveButton->setBounds (644, 280, 72, 24);
+ numVisibleYLabel->setBounds (716, 248, 48, 24);
+ loadTrainingButton->setBounds (564, 316, 72, 24);
+ saveTrainingButton->setBounds (644, 316, 72, 24);
+ clearTrainingButton->setBounds (804, 316, 72, 24);
+ removeTrainingButton->setBounds (724, 316, 72, 24);
+ numGibbsSlider->setBounds (896, 240, 220, 24);
+ rbmDoRaoBlackwellToggleButton->setBounds (556, 60, 112, 24);
+ rbmDoSampleVisibleToggleButton->setBounds (680, 60, 128, 24);
+ lambdaLabel->setBounds (652, 152, 72, 24);
+ sigmaLabel->setBounds (564, 152, 72, 24);
+ rbmUseVisibleGaussianToggleButton->setBounds (680, 92, 156, 24);
+ rbmDoSparseToggleButton->setBounds (556, 92, 96, 24);
+ sparsityLabel->setBounds (740, 152, 72, 24);
+ sigmaDecayLabel->setBounds (836, 152, 72, 24);
+ weightDecayLabel->setBounds (836, 200, 72, 24);
+ m_progressBarSlider->setBounds (828, 20, 220, 24);
+ momentumLabel->setBounds (740, 200, 72, 24);
+ sparsityLearningRateLabel->setBounds (948, 152, 72, 24);
+ weightInitLabel->setBounds (948, 200, 72, 24);
+ rbmLearnVarianceButton->setBounds (1004, 92, 128, 24);
+ rbmNormalizeDataToggleButton->setBounds (1004, 60, 124, 24);
+ m_rbmSelect->setBounds (812, 280, 62, 24);
+ rbmDoSampleBatch->setBounds (844, 60, 128, 24);
+ sizeMiniBatch->setBounds (1064, 20, 48, 24);
+ rbmUseHiddenGaussianToggleButton->setBounds (840, 92, 156, 24);
+ //[UserResized] Add your own custom resize handling here..
+ //[/UserResized]
+}
+
+void MainComponent::buttonClicked (Button* buttonThatWasClicked)
+{
+ //[UserbuttonClicked_Pre]
+ //[/UserbuttonClicked_Pre]
+
+ if (buttonThatWasClicked == trainButton)
+ {
+ //[UserButtonCode_trainButton] -- add your button handler code here..
+ if (isThreadRunning())
+ {
+ m_doStop = true;
+ }
+ else
+ {
+ m_doStop = false;
+ startThread();
+ }
+ //[/UserButtonCode_trainButton]
+ }
+ else if (buttonThatWasClicked == addButton)
+ {
+ //[UserButtonCode_addButton] -- add your button handler code here..
+// m_layers.add(m_pRbmComponent[0]->getTrainingData(), m_weights[0]->getNumVisible());
+// m_pRbmComponent[0]->batchchanged();
+ //[/UserButtonCode_addButton]
+ }
+ else if (buttonThatWasClicked == ShakeButton)
+ {
+ //[UserButtonCode_ShakeButton] -- add your button handler code here..
+ m_pLayer->weightsInit(weightInitLabel->getText().getFloatValue());
+// m_pLayer->batchchanged();
+// m_pLayer->redrawWeights();
+// m_pLayer->redrawReconstruction();
+ //[/UserButtonCode_ShakeButton]
+ }
+ else if (buttonThatWasClicked == createButton)
+ {
+ //[UserButtonCode_createButton] -- add your button handler code here..
+ create("justCreated");
+ //[/UserButtonCode_createButton]
+ }
+ else if (buttonThatWasClicked == loadButton)
+ {
+ //[UserButtonCode_loadButton] -- add your button handler code here..
+ m_stack = new Stack(std::string(projectNameLabel->getText().getCharPointer()));
+ m_stack->load(this);
+ m_rbmSelect->clear(dontSendNotification);
+
+ for (int id=0; id < m_stack->numLayers(); id++)
+ {
+ m_rbmSelect->addItem(String(id), id+1);
+ }
+ m_rbmSelect->setSelectedId(1, sendNotification);
+ //[/UserButtonCode_loadButton]
+ }
+ else if (buttonThatWasClicked == saveButton)
+ {
+ //[UserButtonCode_saveButton] -- add your button handler code here..
+ save();
+ //[/UserButtonCode_saveButton]
+ }
+ else if (buttonThatWasClicked == loadTrainingButton)
+ {
+ //[UserButtonCode_loadTrainingButton] -- add your button handler code here..
+ loadTraining((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
+ m_stack->batchchanged();
+ //[/UserButtonCode_loadTrainingButton]
+ }
+ else if (buttonThatWasClicked == saveTrainingButton)
+ {
+ //[UserButtonCode_saveTrainingButton] -- add your button handler code here..
+ saveTraining((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
+ //[/UserButtonCode_saveTrainingButton]
+ }
+ else if (buttonThatWasClicked == clearTrainingButton)
+ {
+ //[UserButtonCode_clearTrainingButton] -- add your button handler code here..
+ clearTraining();
+ m_stack->batchchanged();
+ //[/UserButtonCode_clearTrainingButton]
+ }
+ else if (buttonThatWasClicked == removeTrainingButton)
+ {
+ //[UserButtonCode_removeTrainingButton] -- add your button handler code here..
+ removeTrainingAt((int)patterSlider->getValue());
+ m_stack->batchchanged();
+ //[/UserButtonCode_removeTrainingButton]
+ }
+ else if (buttonThatWasClicked == rbmDoRaoBlackwellToggleButton)
+ {
+ //[UserButtonCode_rbmDoRaoBlackwellToggleButton] -- add your button handler code here..
+ m_pLayer->params().doRaoBlackwell = buttonThatWasClicked->getToggleState();
+ //[/UserButtonCode_rbmDoRaoBlackwellToggleButton]
+ }
+ else if (buttonThatWasClicked == rbmDoSampleVisibleToggleButton)
+ {
+ //[UserButtonCode_rbmDoSampleVisibleToggleButton] -- add your button handler code here..
+ m_pLayer->params().gibbsDoSampleVisible = buttonThatWasClicked->getToggleState();
+ //[/UserButtonCode_rbmDoSampleVisibleToggleButton]
+ }
+ else if (buttonThatWasClicked == rbmUseVisibleGaussianToggleButton)
+ {
+ //[UserButtonCode_rbmUseVisibleGaussianToggleButton] -- add your button handler code here..
+ //[/UserButtonCode_rbmUseVisibleGaussianToggleButton]
+ }
+ else if (buttonThatWasClicked == rbmDoSparseToggleButton)
+ {
+ //[UserButtonCode_rbmDoSparseToggleButton] -- add your button handler code here..
+ //[/UserButtonCode_rbmDoSparseToggleButton]
+ }
+ else if (buttonThatWasClicked == rbmLearnVarianceButton)
+ {
+ //[UserButtonCode_rbmLearnVarianceButton] -- add your button handler code here..
+ //[/UserButtonCode_rbmLearnVarianceButton]
+ }
+ else if (buttonThatWasClicked == rbmNormalizeDataToggleButton)
+ {
+ //[UserButtonCode_rbmNormalizeDataToggleButton] -- add your button handler code here..
+ //[/UserButtonCode_rbmNormalizeDataToggleButton]
+ }
+ else if (buttonThatWasClicked == rbmDoSampleBatch)
+ {
+ //[UserButtonCode_rbmDoSampleBatch] -- add your button handler code here..
+ m_pLayer->params().doSampleBatch = buttonThatWasClicked->getToggleState();
+ //[/UserButtonCode_rbmDoSampleBatch]
+ }
+ else if (buttonThatWasClicked == rbmUseHiddenGaussianToggleButton)
+ {
+ //[UserButtonCode_rbmUseHiddenGaussianToggleButton] -- add your button handler code here..
+ //[/UserButtonCode_rbmUseHiddenGaussianToggleButton]
+ }
+
+ //[UserbuttonClicked_Post]
+ //[/UserbuttonClicked_Post]
+}
+
+void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
+{
+ //[UsersliderValueChanged_Pre]
+ //[/UsersliderValueChanged_Pre]
+
+ if (sliderThatWasMoved == patterSlider)
+ {
+ //[UserSliderCode_patterSlider] -- add your slider handling code here..
+ m_trainingIndex = (int)sliderThatWasMoved->getValue();
+ //[/UserSliderCode_patterSlider]
+ }
+ else if (sliderThatWasMoved == WeightsSlider)
+ {
+ //[UserSliderCode_WeightsSlider] -- add your slider handling code here..
+ m_weightIndex = (int)sliderThatWasMoved->getValue();
+ //[/UserSliderCode_WeightsSlider]
+ }
+ else if (sliderThatWasMoved == numGibbsSlider)
+ {
+ //[UserSliderCode_numGibbsSlider] -- add your slider handling code here..
+ m_pLayer->params().numGibbs = sliderThatWasMoved->getValue();
+ //[/UserSliderCode_numGibbsSlider]
+ }
+ else if (sliderThatWasMoved == m_progressBarSlider)
+ {
+ //[UserSliderCode_m_progressBarSlider] -- add your slider handling code here..
+ //[/UserSliderCode_m_progressBarSlider]
+ }
+
+ //[UsersliderValueChanged_Post]
+ //[/UsersliderValueChanged_Post]
+}
+
+void MainComponent::labelTextChanged (Label* labelThatHasChanged)
+{
+ //[UserlabelTextChanged_Pre]
+ //[/UserlabelTextChanged_Pre]
+
+ if (labelThatHasChanged == numEpochslabel)
+ {
+ //[UserLabelCode_numEpochslabel] -- add your label text handling code here..
+ //[/UserLabelCode_numEpochslabel]
+ }
+ else if (labelThatHasChanged == learningRateLabel)
+ {
+ //[UserLabelCode_learningRateLabel] -- add your label text handling code here..
+ m_pLayer->params().learningRate = labelThatHasChanged->getText().getFloatValue();
+ //[/UserLabelCode_learningRateLabel]
+ }
+ else if (labelThatHasChanged == numVisibleLabel)
+ {
+ //[UserLabelCode_numVisibleLabel] -- add your label text handling code here..
+ //[/UserLabelCode_numVisibleLabel]
+ }
+ else if (labelThatHasChanged == numHiddenLabel)
+ {
+ //[UserLabelCode_numHiddenLabel] -- add your label text handling code here..
+ //[/UserLabelCode_numHiddenLabel]
+ }
+ else if (labelThatHasChanged == projectNameLabel)
+ {
+ //[UserLabelCode_projectNameLabel] -- add your label text handling code here..
+ //[/UserLabelCode_projectNameLabel]
+ }
+ else if (labelThatHasChanged == numVisibleYLabel)
+ {
+ //[UserLabelCode_numVisibleYLabel] -- add your label text handling code here..
+ //[/UserLabelCode_numVisibleYLabel]
+ }
+ else if (labelThatHasChanged == lambdaLabel)
+ {
+ //[UserLabelCode_lambdaLabel] -- add your label text handling code here..
+ //[/UserLabelCode_lambdaLabel]
+ }
+ else if (labelThatHasChanged == sigmaLabel)
+ {
+ //[UserLabelCode_sigmaLabel] -- add your label text handling code here..
+ //[/UserLabelCode_sigmaLabel]
+ }
+ else if (labelThatHasChanged == sparsityLabel)
+ {
+ //[UserLabelCode_sparsityLabel] -- add your label text handling code here..
+ //[/UserLabelCode_sparsityLabel]
+ }
+ else if (labelThatHasChanged == sigmaDecayLabel)
+ {
+ //[UserLabelCode_sigmaDecayLabel] -- add your label text handling code here..
+ //[/UserLabelCode_sigmaDecayLabel]
+ }
+ else if (labelThatHasChanged == weightDecayLabel)
+ {
+ //[UserLabelCode_weightDecayLabel] -- add your label text handling code here..
+ m_pLayer->params().weightDecay = labelThatHasChanged->getText().getFloatValue();
+ //[/UserLabelCode_weightDecayLabel]
+ }
+ else if (labelThatHasChanged == momentumLabel)
+ {
+ //[UserLabelCode_momentumLabel] -- add your label text handling code here..
+ m_pLayer->params().momentum = labelThatHasChanged->getText().getFloatValue();
+ //[/UserLabelCode_momentumLabel]
+ }
+ else if (labelThatHasChanged == sparsityLearningRateLabel)
+ {
+ //[UserLabelCode_sparsityLearningRateLabel] -- add your label text handling code here..
+ //[/UserLabelCode_sparsityLearningRateLabel]
+ }
+ else if (labelThatHasChanged == weightInitLabel)
+ {
+ //[UserLabelCode_weightInitLabel] -- add your label text handling code here..
+ //[/UserLabelCode_weightInitLabel]
+ }
+ else if (labelThatHasChanged == sizeMiniBatch)
+ {
+ //[UserLabelCode_sizeMiniBatch] -- add your label text handling code here..
+ //[/UserLabelCode_sizeMiniBatch]
+ }
+
+ //[UserlabelTextChanged_Post]
+ //[/UserlabelTextChanged_Post]
+}
+
+void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged)
+{
+ //[UsercomboBoxChanged_Pre]
+ //[/UsercomboBoxChanged_Pre]
+
+ if (comboBoxThatHasChanged == m_rbmSelect)
+ {
+ int index = m_rbmSelect->getSelectedItemIndex();
+ m_pLayer = m_stack->getLayer(index);
+ updateControls();
+ //[/UserComboBoxCode_m_rbmSelect]
+ }
+
+ //[UsercomboBoxChanged_Post]
+ //[/UsercomboBoxChanged_Post]
+}
+
+void MainComponent::mouseMove (const MouseEvent& e)
+{
+ //[UserCode_mouseMove] -- Add your code here...
+ //[/UserCode_mouseMove]
+}
+
+void MainComponent::mouseEnter (const MouseEvent& e)
+{
+ //[UserCode_mouseEnter] -- Add your code here...
+ //[/UserCode_mouseEnter]
+}
+
+void MainComponent::mouseExit (const MouseEvent& e)
+{
+ //[UserCode_mouseExit] -- Add your code here...
+ //[/UserCode_mouseExit]
+}
+
+void MainComponent::mouseDown (const MouseEvent& e)
+{
+ //[UserCode_mouseDown] -- Add your code here...
+ //[/UserCode_mouseDown]
+}
+
+void MainComponent::mouseDrag (const MouseEvent& e)
+{
+ //[UserCode_mouseDrag] -- Add your code here...
+// e.eventComponent->setCentrePosition(e.getPosition().x, e.getPosition().y);
+ std::cout << "Mouse = " << e.x << "," << e.y << std::endl;
+ //[/UserCode_mouseDrag]
+}
+
+void MainComponent::mouseUp (const MouseEvent& e)
+{
+ //[UserCode_mouseUp] -- Add your code here...
+ //[/UserCode_mouseUp]
+}
+
+void MainComponent::mouseDoubleClick (const MouseEvent& e)
+{
+ //[UserCode_mouseDoubleClick] -- Add your code here...
+ //[/UserCode_mouseDoubleClick]
+}
+
+void MainComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel)
+{
+ //[UserCode_mouseWheelMove] -- Add your code here...
+ //[/UserCode_mouseWheelMove]
+}
+
+
+
+//[MiscUserCode] You can add your own definitions of your custom methods or any other code here...
+void MainComponent::save ()
+{
+ m_stack->saveWeights();
+ m_stack->save();
+}
+
+void MainComponent::create(juce::String const &projectName)
+{
+ m_stack = new Stack(std::string(projectName.getCharPointer()));
+
+#if 0
+ int id = m_rbmSelect->getSelectedId()-1;
+ bool shouldAddItem = true;
+ if (id == 0)
+ {
+ m_rbmSelect->clear(dontSendNotification);
+ m_rbmSelect->addItem(String(id), id+1);
+ m_rbmSelect->setSelectedId(id+1, dontSendNotification);
+
+ for (size_t i=0; i < DBN_SIZE; i++)
+ {
+ m_weights[i] = nullptr;
+ m_pRbmComponent[i] = nullptr;
+ }
+ 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_trainingData, *this));
+ }
+ else
+ {
+ shouldAddItem = m_pRbmComponent[id] == nullptr;
+ m_weights[id] = nullptr;
+ m_pRbmComponent[id] = nullptr;
+ 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], *this));
+ }
+ m_pLayer = m_pRbmComponent[id];
+ m_weightsCurr = m_weights[id];
+ m_pLayer->setBounds (16, 140*id+16, 430, 130);
+ m_pLayer->batchchanged();
+ m_pLayer->redrawReconstruction();
+ m_pLayer->redrawWeights();
+
+ if (((id+1) < DBN_SIZE) and shouldAddItem)
+ m_rbmSelect->addItem(String(id+1), id+2);
+#endif
+ updateControls();
+}
+
+void MainComponent::destroy()
+{
+}
+
+const juce::String& MainComponent::getBaseDir()
+{
+ struct passwd *pw = getpwuid(getuid());
+
+ const char *homedir = pw->pw_dir;
+
+
+ m_baseDir = String(homedir) + String("/.rbm/") + projectNameLabel->getText();
+
+ return m_baseDir;
+
+}
+
+void MainComponent::run()
+{
+ m_pLayer->train(m_trainingData);
+}
+
+void MainComponent::onProgressChanged(size_t progressPercent)
+{
+ m_progressBarSlider->setValue(progressPercent);
+}
+
+void MainComponent::updateControls()
+{
+ rbmDoRaoBlackwellToggleButton->setToggleState(m_pLayer->params().doRaoBlackwell, dontSendNotification);
+ rbmDoSampleVisibleToggleButton->setToggleState(m_pLayer->params().gibbsDoSampleVisible, dontSendNotification);
+ rbmDoSampleBatch->setToggleState(m_pLayer->params().doSampleBatch, dontSendNotification);
+
+ weightDecayLabel->setText(String(m_pLayer->params().weightDecay), dontSendNotification);
+ learningRateLabel->setText(String(m_pLayer->params().learningRate), dontSendNotification);
+ momentumLabel->setText(String(m_pLayer->params().momentum), dontSendNotification);
+ numVisibleLabel->setText(String(m_pLayer->numVisibleX()), dontSendNotification );
+ numVisibleYLabel->setText(String(m_pLayer->numVisibleY()), dontSendNotification );
+ numHiddenLabel->setText(String(m_pLayer->numHidden()), dontSendNotification );
+ numEpochslabel->setText(String((int)m_pLayer->params().numEpochs), dontSendNotification);
+ sizeMiniBatch->setText(String((int)m_pLayer->params().miniBatchSize), dontSendNotification);
+
+ numGibbsSlider->setValue(m_pLayer->params().numGibbs, dontSendNotification);
+ WeightsSlider->setRange(0, m_pLayer->numHidden()-1, 1);
+ numGibbsSlider->setValue(m_pLayer->params().numGibbs);
+ WeightsSlider->setValue(m_weightIndex);
+ patterSlider->setValue(m_trainingIndex);
+
+}
diff --git a/source/MainComponent.hpp b/source/MainComponent.hpp
new file mode 100644
index 0000000..fcc2f15
--- /dev/null
+++ b/source/MainComponent.hpp
@@ -0,0 +1,163 @@
+/*
+ ==============================================================================
+
+ This is an automatically generated GUI class created by the Introjucer!
+
+ Be careful when adding custom code to these files, as only the code within
+ the "//[xyz]" and "//[/xyz]" sections will be retained when the file is loaded
+ and re-saved.
+
+ Created with Introjucer version: 3.1.0
+
+ ------------------------------------------------------------------------------
+
+ The Introjucer is part of the JUCE library - "Jules' Utility Class Extensions"
+ Copyright 2004-13 by Raw Material Software Ltd.
+
+ ==============================================================================
+*/
+
+#ifndef __JUCE_HEADER_9002020A4DD09B20__
+#define __JUCE_HEADER_9002020A4DD09B20__
+
+//[Headers] -- You can add your own extra header files here --
+#include
+#include "JuceHeader.h"
+#include "Stack.hpp"
+#include "RbmComponent.hpp"
+
+//[/Headers]
+
+
+
+//==============================================================================
+/**
+ //[Comments]
+ An auto-generated component, created by the Introjucer.
+
+ Describe your class and how it works here!
+ //[/Comments]
+*/
+class MainComponent : public Component,
+ public Thread,
+ public LayerConstructor,
+ public ButtonListener,
+ public SliderListener,
+ public LabelListener,
+ public ComboBoxListener
+{
+public:
+ //==============================================================================
+ MainComponent ();
+ ~MainComponent();
+
+ //==============================================================================
+ //[UserMethods] -- You can add your own custom methods in this section.
+ void onProgressChanged(size_t progressPercent);
+ //[/UserMethods]
+
+ void paint (Graphics& g);
+ void resized();
+ 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);
+ void mouseDown (const MouseEvent& e);
+ void mouseDrag (const MouseEvent& e);
+ void mouseUp (const MouseEvent& e);
+ void mouseDoubleClick (const MouseEvent& e);
+ void mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel);
+
+ Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden)
+ {
+ return new RbmComponent(name, id, numVisibleX, numVisibleY, numHidden);
+ }
+
+private:
+ //[UserVariables] -- You can add your own custom variables in this section.
+ static const size_t DBN_SIZE = 4;
+ ScopedPointer m_stack;
+ Layer *m_pLayer;
+ arma::mat m_trainingData;
+ int m_weightIndex;
+ int m_trainingIndex;
+ void save();
+ void create(juce::String const &projectName);
+ void destroy();
+ const juce::String& getBaseDir();
+ void run();
+ String m_baseDir;
+ TooltipWindow m_toolTipWindow;
+ bool m_doStop;
+ void clearTraining()
+ {
+ }
+
+ void loadTraining(const char *pFilename)
+ {
+ }
+
+ void saveTraining(const char *pFilename)
+ {
+ }
+
+ void removeTrainingAt(size_t index)
+ {
+ }
+
+ void updateControls();
+
+ //[/UserVariables]
+
+ //==============================================================================
+ ScopedPointer trainButton;
+ ScopedPointer addButton;
+ ScopedPointer patterSlider;
+ ScopedPointer ShakeButton;
+ ScopedPointer WeightsSlider;
+ ScopedPointer