- use expectations
- improved gibbs sampling - LayerArray is template class git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@18 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+49
-16
@@ -10,7 +10,7 @@
|
||||
|
||||
#ifndef LAYER_HPP
|
||||
#define LAYER_HPP
|
||||
#include <cstdint>
|
||||
#include <stdint.h>
|
||||
#include "noise.h"
|
||||
#include "Weights.hpp"
|
||||
|
||||
@@ -21,9 +21,8 @@ public:
|
||||
: m_numUnits(numUnits)
|
||||
, m_pProbs(nullptr)
|
||||
, m_pStates(nullptr)
|
||||
, m_pStatesInit(pStatesInit)
|
||||
{
|
||||
setNumUnits(numUnits);
|
||||
setNumUnits(numUnits, pStatesInit);
|
||||
Noise_Init(&m_noise, 0x12345677);
|
||||
}
|
||||
|
||||
@@ -33,7 +32,7 @@ public:
|
||||
Noise_Free(&m_noise);
|
||||
}
|
||||
|
||||
void setNumUnits(uint32_t numUnits)
|
||||
void setNumUnits(uint32_t numUnits, const double *pStatesInit = nullptr)
|
||||
{
|
||||
if (m_numUnits)
|
||||
{
|
||||
@@ -43,23 +42,17 @@ public:
|
||||
m_numUnits = numUnits;
|
||||
if (m_numUnits)
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
m_pProbs = new double[m_numUnits];
|
||||
m_pStates = new double[m_numUnits];
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
probsInit(0);
|
||||
if (pStatesInit)
|
||||
{
|
||||
m_pProbs[i] = 0.0;
|
||||
memcpy(m_pStates, pStatesInit, m_numUnits*sizeof(double));
|
||||
}
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
else
|
||||
{
|
||||
m_pStates[i] = 0.0;
|
||||
}
|
||||
if (m_pStatesInit)
|
||||
{
|
||||
memcpy(m_pStates, m_pStatesInit, m_numUnits*sizeof(double));
|
||||
statesInit(0);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -72,6 +65,37 @@ public:
|
||||
return *this;
|
||||
}
|
||||
|
||||
Layer& operator+= (const Layer &rhs)
|
||||
{
|
||||
uint32_t i;
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pStates[i] += rhs.m_pStates[i];
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
void probsInit(double value) const
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pProbs[i] = value;
|
||||
}
|
||||
}
|
||||
|
||||
void statesInit(double value) const
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pStates[i] = value;
|
||||
}
|
||||
}
|
||||
|
||||
void probsUpdate(const Layer &layer, const Weights &weights) const
|
||||
{
|
||||
uint32_t i;
|
||||
@@ -82,6 +106,16 @@ public:
|
||||
}
|
||||
}
|
||||
|
||||
void statesScale(double kscale) const
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pStates[i] *= kscale;
|
||||
}
|
||||
}
|
||||
|
||||
void statesAssignfromProbs()
|
||||
{
|
||||
memcpy(m_pStates, m_pProbs, m_numUnits*sizeof(double));
|
||||
@@ -128,7 +162,6 @@ private:
|
||||
protected:
|
||||
double *m_pProbs;
|
||||
double *m_pStates;
|
||||
const double *m_pStatesInit;
|
||||
|
||||
virtual double accum(const Layer &layer, const Weights &weights, uint32_t index) const = 0;
|
||||
|
||||
|
||||
+25
-17
@@ -10,27 +10,29 @@
|
||||
|
||||
#ifndef LAYERARRAY_HPP
|
||||
#define LAYERARRAY_HPP
|
||||
#include <cstdint>
|
||||
#include "VisibleLayer.hpp"
|
||||
#include <stdint.h>
|
||||
|
||||
class VisibleLayerArray;
|
||||
template <class T>
|
||||
class LayerArray;
|
||||
|
||||
class VisibleLayerArrayListener
|
||||
template <class T>
|
||||
class LayerArrayListener
|
||||
{
|
||||
public:
|
||||
VisibleLayerArrayListener() {}
|
||||
virtual~VisibleLayerArrayListener() {}
|
||||
LayerArrayListener() {}
|
||||
virtual~LayerArrayListener() {}
|
||||
|
||||
virtual void onChanged(const VisibleLayerArray &obj) = 0;
|
||||
virtual void onChanged(const LayerArray<T> &obj) = 0;
|
||||
};
|
||||
|
||||
class VisibleLayerArray
|
||||
template <class T>
|
||||
class LayerArray
|
||||
{
|
||||
class Entry : public VisibleLayer
|
||||
class Entry : public T
|
||||
{
|
||||
public:
|
||||
Entry(uint32_t numUnits, const double *pInit)
|
||||
: VisibleLayer(numUnits, pInit)
|
||||
: T(numUnits, pInit)
|
||||
, pPrev(nullptr)
|
||||
, pNext(nullptr)
|
||||
{
|
||||
@@ -41,10 +43,11 @@ class VisibleLayerArray
|
||||
Entry *pPrev;
|
||||
Entry *pNext;
|
||||
private:
|
||||
|
||||
};
|
||||
|
||||
public:
|
||||
VisibleLayerArray(VisibleLayerArrayListener *pListener = nullptr)
|
||||
LayerArray(LayerArrayListener<T> *pListener = nullptr)
|
||||
: m_size(0)
|
||||
, m_pRoot(nullptr)
|
||||
, m_ppIndex(nullptr)
|
||||
@@ -52,8 +55,9 @@ public:
|
||||
{
|
||||
}
|
||||
|
||||
virtual ~VisibleLayerArray()
|
||||
virtual ~LayerArray()
|
||||
{
|
||||
m_pListener = nullptr;
|
||||
clear();
|
||||
}
|
||||
|
||||
@@ -82,7 +86,6 @@ public:
|
||||
|
||||
void clear()
|
||||
{
|
||||
// rebuildIndex();
|
||||
if (m_ppIndex)
|
||||
{
|
||||
for (int i=0; i < m_size; i++)
|
||||
@@ -101,11 +104,16 @@ public:
|
||||
m_pListener->onChanged(*this);
|
||||
}
|
||||
|
||||
VisibleLayer& getAt(uint32_t index)
|
||||
T& getAt(uint32_t index)
|
||||
{
|
||||
return *m_ppIndex[index];
|
||||
}
|
||||
|
||||
T& operator[](uint32_t index)
|
||||
{
|
||||
return getAt(index);
|
||||
}
|
||||
|
||||
void removeAt(uint32_t index)
|
||||
{
|
||||
if (m_ppIndex[index]->pPrev)
|
||||
@@ -145,12 +153,12 @@ public:
|
||||
|
||||
for (i=0; i < m_size; i++)
|
||||
{
|
||||
uint32_t numUnits = ((VisibleLayer*)m_ppIndex[i])->getNumUnits();
|
||||
uint32_t numUnits = m_ppIndex[i]->getNumUnits();
|
||||
fprintf(pFile, "%d\n", numUnits);
|
||||
|
||||
for (j=0; j < numUnits; j++)
|
||||
{
|
||||
fprintf(pFile, "%3.6f\n", ((VisibleLayer*)m_ppIndex[i])->getStates()[j]);
|
||||
fprintf(pFile, "%3.6f\n", m_ppIndex[i]->getStates()[j]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -194,7 +202,7 @@ private:
|
||||
uint32_t m_size;
|
||||
Entry *m_pRoot;
|
||||
Entry **m_ppIndex;
|
||||
VisibleLayerArrayListener *m_pListener;
|
||||
LayerArrayListener<T> *m_pListener;
|
||||
void allocIndex(size_t size)
|
||||
{
|
||||
if (m_ppIndex)
|
||||
|
||||
+107
-32
@@ -18,13 +18,16 @@
|
||||
*/
|
||||
|
||||
//[Headers] You can add your own extra header files here...
|
||||
#ifdef WIN32
|
||||
#include <Windows.h>
|
||||
#endif
|
||||
//[/Headers]
|
||||
|
||||
#include "MainComponent.h"
|
||||
|
||||
|
||||
//[MiscUserDefs] You can add your own user definitions and misc code here...
|
||||
#ifdef WIN32
|
||||
void mylog(const char* format, ...)
|
||||
{
|
||||
char log_buf[257];
|
||||
@@ -34,6 +37,17 @@ void mylog(const char* format, ...)
|
||||
va_end(argptr);
|
||||
OutputDebugStringA (log_buf);
|
||||
}
|
||||
#else
|
||||
void mylog(const char* format, ...)
|
||||
{
|
||||
char log_buf[257];
|
||||
va_list argptr;
|
||||
va_start(argptr, format);
|
||||
vprintf(format, argptr);
|
||||
va_end(argptr);
|
||||
}
|
||||
#endif
|
||||
|
||||
//[/MiscUserDefs]
|
||||
|
||||
//==============================================================================
|
||||
@@ -53,7 +67,7 @@ MainComponent::MainComponent ()
|
||||
addButton->addListener (this);
|
||||
|
||||
addAndMakeVisible (patterSlider = new Slider ("Pattern slider"));
|
||||
patterSlider->setRange (0, 10, 1);
|
||||
patterSlider->setRange (0, 0, 1);
|
||||
patterSlider->setSliderStyle (Slider::LinearHorizontal);
|
||||
patterSlider->setTextBoxStyle (Slider::TextBoxLeft, false, 80, 20);
|
||||
patterSlider->addListener (this);
|
||||
@@ -67,7 +81,7 @@ MainComponent::MainComponent ()
|
||||
ShakeButton->addListener (this);
|
||||
|
||||
addAndMakeVisible (WeightsSlider = new Slider ("Weights slider"));
|
||||
WeightsSlider->setRange (0, 10, 1);
|
||||
WeightsSlider->setRange (0, 0, 1);
|
||||
WeightsSlider->setSliderStyle (Slider::LinearHorizontal);
|
||||
WeightsSlider->setTextBoxStyle (Slider::TextBoxLeft, false, 80, 20);
|
||||
WeightsSlider->addListener (this);
|
||||
@@ -158,6 +172,20 @@ MainComponent::MainComponent ()
|
||||
removeTrainingButton->setButtonText (TRANS("Remove T"));
|
||||
removeTrainingButton->addListener (this);
|
||||
|
||||
addAndMakeVisible (numGibbsSlider = new Slider ("Num. Gibbs slider"));
|
||||
numGibbsSlider->setRange (1, 10, 1);
|
||||
numGibbsSlider->setSliderStyle (Slider::LinearHorizontal);
|
||||
numGibbsSlider->setTextBoxStyle (Slider::TextBoxLeft, false, 80, 20);
|
||||
numGibbsSlider->addListener (this);
|
||||
|
||||
addAndMakeVisible (reconstructEquButton = new TextButton ("Reconstruct Equilibrium button"));
|
||||
reconstructEquButton->setButtonText (TRANS("Reconst Equ."));
|
||||
reconstructEquButton->addListener (this);
|
||||
|
||||
addAndMakeVisible (rbmUseExpectationsToggleButton = new ToggleButton ("rbmUseExpectations toggle button"));
|
||||
rbmUseExpectationsToggleButton->setButtonText (TRANS("Use Expectations"));
|
||||
rbmUseExpectationsToggleButton->addListener (this);
|
||||
|
||||
|
||||
//[UserPreSize]
|
||||
m_vNumX = 16;
|
||||
@@ -166,11 +194,8 @@ MainComponent::MainComponent ()
|
||||
m_vNumX_next = 16;
|
||||
m_vNumY_next = 16;
|
||||
m_hNum_next = 64;
|
||||
|
||||
// m_pWeights = new Weights(m_vNumX*m_vNumY, m_hNum);
|
||||
// m_pWeights = new Weights("C:\\Dokumente und Einstellungen\\Jens\\Desktop\\weights.dat");
|
||||
m_weights.load("C:\\Dokumente und Einstellungen\\Jens\\Desktop\\weights.dat");
|
||||
m_layers.load("C:\\Dokumente und Einstellungen\\Jens\\Desktop\\training_states.dat");
|
||||
m_numGibbs = 1;
|
||||
m_weights.setUnits(m_vNumX*m_vNumY, m_hNum);
|
||||
create();
|
||||
|
||||
//[/UserPreSize]
|
||||
@@ -181,11 +206,9 @@ MainComponent::MainComponent ()
|
||||
//[Constructor] You can add your own custom stuff here..
|
||||
|
||||
projectNameLabel->setText(String("TestPrj"), dontSendNotification );
|
||||
numVisibleLabel->setText(String(m_vNumX), dontSendNotification );
|
||||
numVisibleYLabel->setText(String(m_vNumY), dontSendNotification );
|
||||
numHiddenLabel->setText(String(m_hNum), dontSendNotification );
|
||||
numEpochslabel->setText(String(1000), dontSendNotification );
|
||||
numEpochslabel->setText(String(100), dontSendNotification );
|
||||
learningRateLabel->setText(String(0.2), dontSendNotification );
|
||||
rbmUseExpectationsToggleButton->setToggleState(false, true);
|
||||
//[/Constructor]
|
||||
}
|
||||
|
||||
@@ -214,6 +237,9 @@ MainComponent::~MainComponent()
|
||||
saveTrainingButton = nullptr;
|
||||
clearTrainingButton = nullptr;
|
||||
removeTrainingButton = nullptr;
|
||||
numGibbsSlider = nullptr;
|
||||
reconstructEquButton = nullptr;
|
||||
rbmUseExpectationsToggleButton = nullptr;
|
||||
|
||||
|
||||
//[Destructor]. You can add your own custom destruction code here..
|
||||
@@ -245,8 +271,8 @@ void MainComponent::resized()
|
||||
reconstructButton->setBounds (24, 280, 72, 24);
|
||||
ShakeButton->setBounds (120, 320, 72, 24);
|
||||
WeightsSlider->setBounds (224, 320, 184, 24);
|
||||
numEpochslabel->setBounds (24, 248, 72, 24);
|
||||
learningRateLabel->setBounds (120, 248, 72, 24);
|
||||
numEpochslabel->setBounds (24, 240, 72, 24);
|
||||
learningRateLabel->setBounds (120, 240, 72, 24);
|
||||
testButton->setBounds (120, 176, 72, 24);
|
||||
numVisibleLabel->setBounds (440, 256, 72, 24);
|
||||
numHiddenLabel->setBounds (488, 288, 72, 24);
|
||||
@@ -259,6 +285,9 @@ void MainComponent::resized()
|
||||
saveTrainingButton->setBounds (520, 48, 72, 24);
|
||||
clearTrainingButton->setBounds (520, 88, 72, 24);
|
||||
removeTrainingButton->setBounds (440, 88, 72, 24);
|
||||
numGibbsSlider->setBounds (224, 240, 184, 24);
|
||||
reconstructEquButton->setBounds (24, 320, 72, 24);
|
||||
rbmUseExpectationsToggleButton->setBounds (224, 176, 150, 24);
|
||||
//[UserResized] Add your own custom resize handling here..
|
||||
Draw->setBounds (16, 16, 100, 100);
|
||||
Draw2->setBounds (110+16, 16, 100, 100);
|
||||
@@ -274,17 +303,13 @@ 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_pRbm->train(m_layers, numEpochslabel->getText().getIntValue(), learningRateLabel->getText().getFloatValue(), m_numGibbs, m_rbmUseExpectations);
|
||||
//[/UserButtonCode_trainButton]
|
||||
}
|
||||
else if (buttonThatWasClicked == addButton)
|
||||
{
|
||||
//[UserButtonCode_addButton] -- add your button handler code here..
|
||||
Draw2->setData(Draw->getData());
|
||||
for (int i=0; i < m_vNumX*m_vNumY; i++)
|
||||
{
|
||||
mylog("%3.6f\n", Draw->getData()[i]);
|
||||
}
|
||||
m_layers.add(Draw->getData(), m_vNumX*m_vNumY);
|
||||
patterSlider->setRange (0, m_layers.getSize()-1, 1);
|
||||
//[/UserButtonCode_addButton]
|
||||
@@ -357,6 +382,27 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
||||
m_layers.removeAt((int)patterSlider->getValue());
|
||||
//[/UserButtonCode_removeTrainingButton]
|
||||
}
|
||||
else if (buttonThatWasClicked == reconstructEquButton)
|
||||
{
|
||||
//[UserButtonCode_reconstructEquButton] -- add your button handler code here..
|
||||
uint32_t i;
|
||||
const double *pV, *pH;
|
||||
|
||||
pV = Draw->getData();
|
||||
for (i=0; i < 1000; i++)
|
||||
{
|
||||
pH = m_pRbm->toHidden(pV);
|
||||
pV = m_pRbm->toVisible(pH);
|
||||
Draw2->setData(pV);
|
||||
}
|
||||
//[/UserButtonCode_reconstructEquButton]
|
||||
}
|
||||
else if (buttonThatWasClicked == rbmUseExpectationsToggleButton)
|
||||
{
|
||||
//[UserButtonCode_rbmUseExpectationsToggleButton] -- add your button handler code here..
|
||||
m_rbmUseExpectations = buttonThatWasClicked->getToggleState();
|
||||
//[/UserButtonCode_rbmUseExpectationsToggleButton]
|
||||
}
|
||||
|
||||
//[UserbuttonClicked_Post]
|
||||
//[/UserbuttonClicked_Post]
|
||||
@@ -372,7 +418,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
|
||||
//[UserSliderCode_patterSlider] -- add your slider handling code here..
|
||||
if (m_layers.getSize() > 0)
|
||||
{
|
||||
VisibleLayer &p = m_layers.getAt((int)sliderThatWasMoved->getValue());
|
||||
VisibleLayer &p = (VisibleLayer&)m_layers.getAt((int)sliderThatWasMoved->getValue());
|
||||
Draw2->setData(p.getStates());
|
||||
}
|
||||
//[/UserSliderCode_patterSlider]
|
||||
@@ -380,8 +426,12 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
|
||||
else if (sliderThatWasMoved == WeightsSlider)
|
||||
{
|
||||
//[UserSliderCode_WeightsSlider] -- add your slider handling code here..
|
||||
double *pW = m_weights.getWeights()[(int)sliderThatWasMoved->getValue()];
|
||||
ScopedPointer<double> pTemp = new double [m_vNumX*m_vNumY];
|
||||
double **ppW = m_weights.getWeights();
|
||||
if (!ppW)
|
||||
return;
|
||||
|
||||
double *pW = ppW[(int)sliderThatWasMoved->getValue()];
|
||||
double *pTemp = new double [m_vNumX*m_vNumY];
|
||||
double min = +1E12;
|
||||
double max = -1E12;
|
||||
|
||||
@@ -401,8 +451,15 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
|
||||
}
|
||||
|
||||
DrawWeights->setData(pTemp);
|
||||
delete pTemp;
|
||||
//[/UserSliderCode_WeightsSlider]
|
||||
}
|
||||
else if (sliderThatWasMoved == numGibbsSlider)
|
||||
{
|
||||
//[UserSliderCode_numGibbsSlider] -- add your slider handling code here..
|
||||
m_numGibbs = (uint32_t)sliderThatWasMoved->getValue();
|
||||
//[/UserSliderCode_numGibbsSlider]
|
||||
}
|
||||
|
||||
//[UsersliderValueChanged_Post]
|
||||
//[/UsersliderValueChanged_Post]
|
||||
@@ -457,8 +514,6 @@ void MainComponent::labelTextChanged (Label* labelThatHasChanged)
|
||||
void MainComponent::load ()
|
||||
{
|
||||
m_weights.load((String(getBaseDir() + String(".weights.dat"))).toUTF8());
|
||||
m_layers.clear();
|
||||
m_layers.load((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
|
||||
|
||||
m_vNumX = m_vNumY = (int)sqrt((float)m_weights.getNumVisible());
|
||||
m_hNum = m_weights.getNumHidden();
|
||||
@@ -479,7 +534,10 @@ void MainComponent::create()
|
||||
addAndMakeVisible (Draw = new DrawComponent (m_vNumX, m_vNumX));
|
||||
addAndMakeVisible (Draw2 = new DrawComponent (m_vNumX, m_vNumX));
|
||||
addAndMakeVisible (DrawWeights = new DrawComponent (m_vNumX, m_vNumX));
|
||||
m_pRbm = new Rbm(m_weights);
|
||||
numVisibleLabel->setText(String(m_vNumX), dontSendNotification );
|
||||
numVisibleYLabel->setText(String(m_vNumY), dontSendNotification );
|
||||
numHiddenLabel->setText(String(m_hNum), dontSendNotification );
|
||||
m_pRbm = new Rbm(m_weights, this);
|
||||
WeightsSlider->setRange(0, m_hNum-1, 1);
|
||||
resized();
|
||||
}
|
||||
@@ -490,17 +548,23 @@ void MainComponent::destroy()
|
||||
|
||||
const juce::String& MainComponent::getBaseDir()
|
||||
{
|
||||
static String baseDir(String("C:\\Dokumente und Einstellungen\\Jens\\Desktop\\") + projectNameLabel->getText());
|
||||
m_baseDir = String("./") + projectNameLabel->getText();
|
||||
|
||||
return baseDir;
|
||||
return m_baseDir;
|
||||
|
||||
}
|
||||
|
||||
void MainComponent::onChanged(const VisibleLayerArray &obj)
|
||||
void MainComponent::onChanged(const LayerArray<VisibleLayer> &obj)
|
||||
{
|
||||
patterSlider->setRange(0, std::max(0,(int)m_layers.getSize()-1), 1);
|
||||
}
|
||||
|
||||
void MainComponent::onEpochTrained(const Rbm &obj)
|
||||
{
|
||||
m_trainingProgress = obj.getProgress();
|
||||
mylog("Training %f %%\n", m_trainingProgress*100);
|
||||
}
|
||||
|
||||
//[/MiscUserCode]
|
||||
|
||||
|
||||
@@ -514,7 +578,7 @@ void MainComponent::onChanged(const VisibleLayerArray &obj)
|
||||
BEGIN_JUCER_METADATA
|
||||
|
||||
<JUCER_COMPONENT documentType="Component" className="MainComponent" componentName=""
|
||||
parentClasses="public Component, public VisibleLayerArrayListener"
|
||||
parentClasses="public Component, public LayerArrayListener<VisibleLayer>, public RbmListener"
|
||||
constructorParams="" variableInitialisers="m_layers(this), m_pRbm(nullptr) Draw(nullptr), Draw2(nullptr), DrawWeights(nullptr)"
|
||||
snapPixels="8" snapActive="1" snapShown="1" overlayOpacity="0.330"
|
||||
fixedSize="1" initialWidth="600" initialHeight="400">
|
||||
@@ -527,7 +591,7 @@ BEGIN_JUCER_METADATA
|
||||
connectedEdges="0" needsCallback="1" radioGroupId="0"/>
|
||||
<SLIDER name="Pattern slider" id="c3e0a2c816db81d1" memberName="patterSlider"
|
||||
virtualName="" explicitFocusOrder="0" pos="224 280 184 24" min="0"
|
||||
max="10" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft"
|
||||
max="0" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft"
|
||||
textBoxEditable="1" textBoxWidth="80" textBoxHeight="20" skewFactor="1"/>
|
||||
<TEXTBUTTON name="Reconstruct button" id="c1901d121a5819d6" memberName="reconstructButton"
|
||||
virtualName="" explicitFocusOrder="0" pos="24 280 72 24" buttonText="Reconstruct"
|
||||
@@ -537,15 +601,15 @@ BEGIN_JUCER_METADATA
|
||||
connectedEdges="0" needsCallback="1" radioGroupId="0"/>
|
||||
<SLIDER name="Weights slider" id="699075a1fc01458a" memberName="WeightsSlider"
|
||||
virtualName="" explicitFocusOrder="0" pos="224 320 184 24" min="0"
|
||||
max="10" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft"
|
||||
max="0" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft"
|
||||
textBoxEditable="1" textBoxWidth="80" textBoxHeight="20" skewFactor="1"/>
|
||||
<LABEL name="Num Epochs label" id="b23ae372ee931474" memberName="numEpochslabel"
|
||||
virtualName="" explicitFocusOrder="0" pos="24 248 72 24" edTextCol="ff000000"
|
||||
virtualName="" explicitFocusOrder="0" pos="24 240 72 24" edTextCol="ff000000"
|
||||
edBkgCol="0" labelText="99999" editableSingleClick="1" editableDoubleClick="1"
|
||||
focusDiscardsChanges="0" fontname="Default font" fontsize="15"
|
||||
bold="0" italic="0" justification="36"/>
|
||||
<LABEL name="Learning Rate label" id="49611a27914e910d" memberName="learningRateLabel"
|
||||
virtualName="" explicitFocusOrder="0" pos="120 248 72 24" edTextCol="ff000000"
|
||||
virtualName="" explicitFocusOrder="0" pos="120 240 72 24" edTextCol="ff000000"
|
||||
edBkgCol="0" labelText="0.001" editableSingleClick="1" editableDoubleClick="1"
|
||||
focusDiscardsChanges="0" fontname="Default font" fontsize="15"
|
||||
bold="0" italic="0" justification="36"/>
|
||||
@@ -593,6 +657,17 @@ BEGIN_JUCER_METADATA
|
||||
<TEXTBUTTON name="Remove Training button" id="f0e1d068c39986f3" memberName="removeTrainingButton"
|
||||
virtualName="" explicitFocusOrder="0" pos="440 88 72 24" buttonText="Remove T"
|
||||
connectedEdges="0" needsCallback="1" radioGroupId="0"/>
|
||||
<SLIDER name="Num. Gibbs slider" id="fdd9d5ca4e05e18c" memberName="numGibbsSlider"
|
||||
virtualName="" explicitFocusOrder="0" pos="224 240 184 24" min="1"
|
||||
max="10" int="1" style="LinearHorizontal" textBoxPos="TextBoxLeft"
|
||||
textBoxEditable="1" textBoxWidth="80" textBoxHeight="20" skewFactor="1"/>
|
||||
<TEXTBUTTON name="Reconstruct Equilibrium button" id="208a587f5556207" memberName="reconstructEquButton"
|
||||
virtualName="" explicitFocusOrder="0" pos="24 320 72 24" buttonText="Reconst Equ."
|
||||
connectedEdges="0" needsCallback="1" radioGroupId="0"/>
|
||||
<TOGGLEBUTTON name="rbmUseExpectations toggle button" id="62884e37cd027719"
|
||||
memberName="rbmUseExpectationsToggleButton" virtualName="" explicitFocusOrder="0"
|
||||
pos="224 176 150 24" buttonText="Use Expectations" connectedEdges="0"
|
||||
needsCallback="1" radioGroupId="0" state="0"/>
|
||||
</JUCER_COMPONENT>
|
||||
|
||||
END_JUCER_METADATA
|
||||
|
||||
+12
-4
@@ -38,7 +38,8 @@
|
||||
//[/Comments]
|
||||
*/
|
||||
class MainComponent : public Component,
|
||||
public VisibleLayerArrayListener,
|
||||
public LayerArrayListener<VisibleLayer>,
|
||||
public RbmListener,
|
||||
public ButtonListener,
|
||||
public SliderListener,
|
||||
public LabelListener
|
||||
@@ -67,20 +68,24 @@ private:
|
||||
ScopedPointer<DrawComponent> Draw;
|
||||
ScopedPointer<DrawComponent> Draw2;
|
||||
ScopedPointer<DrawComponent> DrawWeights;
|
||||
Array<VisibleLayer>m_trainingData;
|
||||
uint32_t m_vNumX;
|
||||
uint32_t m_vNumY;
|
||||
uint32_t m_hNum;
|
||||
uint32_t m_vNumX_next;
|
||||
uint32_t m_vNumY_next;
|
||||
uint32_t m_hNum_next;
|
||||
VisibleLayerArray m_layers;
|
||||
LayerArray<VisibleLayer> m_layers;
|
||||
void load();
|
||||
void save();
|
||||
void create();
|
||||
void destroy();
|
||||
const juce::String& getBaseDir();
|
||||
void onChanged(const VisibleLayerArray &obj);
|
||||
void onChanged(const LayerArray<VisibleLayer> &obj);
|
||||
void onEpochTrained(const Rbm &obj);
|
||||
String m_baseDir;
|
||||
double m_trainingProgress;
|
||||
uint32_t m_numGibbs;
|
||||
bool m_rbmUseExpectations;
|
||||
//[/UserVariables]
|
||||
|
||||
//==============================================================================
|
||||
@@ -104,6 +109,9 @@ private:
|
||||
ScopedPointer<TextButton> saveTrainingButton;
|
||||
ScopedPointer<TextButton> clearTrainingButton;
|
||||
ScopedPointer<TextButton> removeTrainingButton;
|
||||
ScopedPointer<Slider> numGibbsSlider;
|
||||
ScopedPointer<TextButton> reconstructEquButton;
|
||||
ScopedPointer<ToggleButton> rbmUseExpectationsToggleButton;
|
||||
|
||||
|
||||
//==============================================================================
|
||||
|
||||
+80
-60
@@ -16,13 +16,26 @@
|
||||
void mylog(const char* format, ...);
|
||||
#define printf mylog
|
||||
|
||||
class Rbm;
|
||||
|
||||
class RbmListener
|
||||
{
|
||||
public:
|
||||
RbmListener() {}
|
||||
virtual ~RbmListener() {}
|
||||
|
||||
virtual void onEpochTrained(const Rbm &obj) = 0;
|
||||
};
|
||||
|
||||
class Rbm
|
||||
{
|
||||
public:
|
||||
Rbm(Weights &weights)
|
||||
Rbm(Weights &weights, RbmListener *pListener = nullptr)
|
||||
: m_w(weights)
|
||||
, m_pListener(pListener)
|
||||
, m_tv(weights.getNumVisible())
|
||||
, m_th(weights.getNumHidden())
|
||||
, m_progress(0)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
}
|
||||
@@ -32,95 +45,100 @@ public:
|
||||
Noise_Free(&m_noise);
|
||||
}
|
||||
|
||||
void weightsUpdate(VisibleLayer &v, VisibleLayer &vr, HiddenLayer &h, HiddenLayer &hr, double mu)
|
||||
void weightsUpdate(VisibleLayer &v, HiddenLayer &h, double mu)
|
||||
{
|
||||
uint32_t i, j;
|
||||
double dw;
|
||||
|
||||
double **ppW = m_w.getWeights();
|
||||
const double *pH = h.getStates();
|
||||
const double *pV = v.getStates();
|
||||
|
||||
// Update weights
|
||||
for (i=0; i < m_w.getNumHidden(); i++)
|
||||
{
|
||||
dw = 0;
|
||||
for (j=0; j < m_w.getNumVisible(); j++)
|
||||
{
|
||||
dw = v.getStates()[j] * h.getStates()[i];
|
||||
m_w.getWeights()[i][j] += mu*dw;
|
||||
dw = pV[j] * pH[i];
|
||||
ppW[i][j] += mu*dw;
|
||||
}
|
||||
}
|
||||
|
||||
for (i=0; i < m_w.getNumHidden(); i++)
|
||||
{
|
||||
dw = 0;
|
||||
for (j=0; j < m_w.getNumVisible(); j++)
|
||||
{
|
||||
dw = vr.getStates()[j] * hr.getStates()[i];
|
||||
m_w.getWeights()[i][j] -= mu*dw;
|
||||
}
|
||||
}
|
||||
#if 1
|
||||
double *pBias = m_w.getBiasVisible();
|
||||
for (i=0; i < m_w.getNumVisible(); i++)
|
||||
{
|
||||
dw = v.getStates()[i] - vr.getStates()[i];
|
||||
m_w.getBiasVisible()[i] += mu*dw;
|
||||
dw = pV[i];
|
||||
pBias[i] += mu*dw;
|
||||
}
|
||||
|
||||
pBias = m_w.getBiasHidden();
|
||||
for (i=0; i < m_w.getNumHidden(); i++)
|
||||
{
|
||||
dw = h.getStates()[i] - hr.getStates()[i];
|
||||
m_w.getBiasHidden()[i] += mu*dw;
|
||||
dw = pH[i];
|
||||
pBias[i] += mu*dw;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
void train(VisibleLayerArray &vts, uint32_t numEpochs, double mu)
|
||||
void train(LayerArray<VisibleLayer> &vts, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, bool useExpectations=false)
|
||||
{
|
||||
uint32_t t;
|
||||
uint32_t epoch;
|
||||
uint32_t trainingPatternIndex;
|
||||
VisibleLayer vr(m_w.getNumVisible());
|
||||
uint32_t gibbs;
|
||||
VisibleLayer v(m_w.getNumVisible());
|
||||
HiddenLayer h(m_w.getNumHidden());
|
||||
HiddenLayer hr(m_w.getNumHidden());
|
||||
const uint32_t monitorInterval = 100; // epochs
|
||||
uint32_t monitorCount = monitorInterval; // epochs
|
||||
Weights w = m_w;
|
||||
|
||||
double dProgress = 1.0/numEpochs;
|
||||
m_progress = 0;
|
||||
|
||||
for (epoch=0; epoch < numEpochs; epoch++)
|
||||
{
|
||||
trainingPatternIndex = (uint32_t)((vts.getSize())*Noise_Uniform(&m_noise, 0.5));
|
||||
|
||||
if (trainingPatternIndex == vts.getSize())
|
||||
continue;
|
||||
|
||||
// Assign training data
|
||||
VisibleLayer &vt = vts.getAt(trainingPatternIndex);
|
||||
|
||||
// Create hidden layer base on training data
|
||||
h.probsUpdate(vt, m_w);
|
||||
// h.statesAssignfromProbs();
|
||||
h.statesUpdateStochastic();
|
||||
|
||||
// Create visible reconstruction (a fantasy...)
|
||||
vr = vt;
|
||||
vr.probsUpdate(h, m_w);
|
||||
// vr.statesAssignfromProbs();
|
||||
vr.statesUpdateStochastic();
|
||||
|
||||
// Create hidden reconstruction
|
||||
hr.probsUpdate(vr, m_w);
|
||||
// hr.statesAssignfromProbs();
|
||||
hr.statesUpdateStochastic();
|
||||
|
||||
// Update weights
|
||||
weightsUpdate(vt, vr, h, hr, mu);
|
||||
|
||||
if (!monitorCount)
|
||||
for (t=0; t < vts.getSize(); t++)
|
||||
{
|
||||
monitorCount = monitorInterval;
|
||||
// printf("Epoch #%d\n", epoch);
|
||||
// prob();
|
||||
// Create hidden layer base on training data
|
||||
h.probsUpdate(vts[t], w);
|
||||
h.statesUpdateStochastic();
|
||||
|
||||
// Update weights (positive phase)
|
||||
weightsUpdate(vts[t], h, +mu/vts.getSize());
|
||||
|
||||
for (gibbs=0; gibbs < numGibbs; gibbs++)
|
||||
{
|
||||
// Create visible reconstruction (a fantasy...)
|
||||
v.probsUpdate(h, w);
|
||||
v.statesUpdateStochastic();
|
||||
|
||||
// Create hidden reconstruction
|
||||
h.probsUpdate(v, w);
|
||||
h.statesUpdateStochastic();
|
||||
}
|
||||
// Update weights (negative phase)
|
||||
h.statesAssignfromProbs();
|
||||
weightsUpdate(v, h, -mu/vts.getSize());
|
||||
|
||||
if (!useExpectations)
|
||||
{
|
||||
w = m_w;
|
||||
}
|
||||
|
||||
}
|
||||
if (useExpectations)
|
||||
{
|
||||
w = m_w;
|
||||
}
|
||||
|
||||
m_progress += dProgress;
|
||||
if (m_pListener)
|
||||
{
|
||||
m_pListener->onEpochTrained(*this);
|
||||
}
|
||||
monitorCount--;
|
||||
}
|
||||
}
|
||||
|
||||
double getProgress() const
|
||||
{
|
||||
return m_progress;
|
||||
}
|
||||
|
||||
double getEnergy(VisibleLayer &v, HiddenLayer &h)
|
||||
{
|
||||
uint32_t i, j;
|
||||
@@ -138,7 +156,7 @@ public:
|
||||
return energy;
|
||||
}
|
||||
|
||||
void prob(VisibleLayerArray &vts)
|
||||
void prob(LayerArray<VisibleLayer> &vts)
|
||||
{
|
||||
uint32_t i, j;
|
||||
double z;
|
||||
@@ -264,9 +282,11 @@ public:
|
||||
|
||||
private:
|
||||
Weights &m_w;
|
||||
RbmListener *m_pListener;
|
||||
VisibleLayer m_tv;
|
||||
HiddenLayer m_th;
|
||||
noise_gen_t m_noise;
|
||||
double m_progress;
|
||||
|
||||
};
|
||||
|
||||
|
||||
+55
-15
@@ -10,7 +10,7 @@
|
||||
|
||||
#ifndef WEIGHTS_HPP
|
||||
#define WEIGHTS_HPP
|
||||
#include <cstdint>
|
||||
#include <stdint.h>
|
||||
#include "noise.h"
|
||||
|
||||
class Weights
|
||||
@@ -33,9 +33,8 @@ public:
|
||||
, m_numVisible(numVisible)
|
||||
, m_numHidden(numHidden)
|
||||
{
|
||||
alloc();
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
|
||||
alloc(numVisible, numHidden);
|
||||
shuffle(0);
|
||||
}
|
||||
|
||||
@@ -49,6 +48,18 @@ public:
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
}
|
||||
|
||||
Weights(const Weights &src)
|
||||
: m_ppW(nullptr)
|
||||
, m_pBiasVisible(nullptr)
|
||||
, m_pBiasHidden(nullptr)
|
||||
, m_numVisible(0)
|
||||
, m_numHidden(0)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
alloc(src.m_numVisible, src.m_numHidden);
|
||||
*this = src;
|
||||
}
|
||||
|
||||
~Weights()
|
||||
{
|
||||
Noise_Free(&m_noise);
|
||||
@@ -57,9 +68,7 @@ public:
|
||||
|
||||
void setUnits(uint32_t numVisible, uint32_t numHidden)
|
||||
{
|
||||
m_numVisible = numVisible;
|
||||
m_numHidden = numHidden;
|
||||
alloc();
|
||||
alloc(numVisible, numHidden);
|
||||
shuffle(0);
|
||||
}
|
||||
|
||||
@@ -87,6 +96,30 @@ public:
|
||||
}
|
||||
}
|
||||
|
||||
Weights& operator= (const Weights &rhs)
|
||||
{
|
||||
uint32_t i, j;
|
||||
|
||||
for (j=0; j < m_numVisible; j++)
|
||||
{
|
||||
m_pBiasVisible[j] = rhs.m_pBiasVisible[j];
|
||||
}
|
||||
|
||||
for (i=0; i < m_numHidden; i++)
|
||||
{
|
||||
m_pBiasHidden[i] = rhs.m_pBiasHidden[i];
|
||||
}
|
||||
|
||||
for (i=0; i < m_numHidden; i++)
|
||||
{
|
||||
for (j=0; j < m_numVisible; j++)
|
||||
{
|
||||
m_ppW[i][j] = rhs.m_ppW[i][j];
|
||||
}
|
||||
}
|
||||
return *this;
|
||||
}
|
||||
|
||||
double **getWeights() const
|
||||
{
|
||||
return m_ppW;
|
||||
@@ -183,6 +216,8 @@ public:
|
||||
|
||||
void load(const char *pFilename)
|
||||
{
|
||||
uint32_t numVisible;
|
||||
uint32_t numHidden;
|
||||
FILE *pFile;
|
||||
|
||||
pFile = fopen(pFilename,"r");
|
||||
@@ -190,26 +225,25 @@ public:
|
||||
if (!pFile)
|
||||
return;
|
||||
|
||||
m_numVisible = m_numHidden = 0;
|
||||
fscanf(pFile, "%d %d\n", &m_numVisible, &m_numHidden);
|
||||
fscanf(pFile, "%d %d\n", &numVisible, &numHidden);
|
||||
|
||||
alloc();
|
||||
alloc(numVisible, numHidden);
|
||||
|
||||
uint32_t i, j;
|
||||
float v;
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
for (i=0; i < numVisible; i++)
|
||||
{
|
||||
fscanf(pFile, "%f", &v);
|
||||
m_pBiasVisible[i] = v;
|
||||
}
|
||||
for (i=0; i < m_numHidden; i++)
|
||||
for (i=0; i < numHidden; i++)
|
||||
{
|
||||
fscanf(pFile, "%f", &v);
|
||||
m_pBiasHidden[i] = v;
|
||||
}
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
for (i=0; i < numVisible; i++)
|
||||
{
|
||||
for (j=0; j < m_numHidden; j++)
|
||||
for (j=0; j < numHidden; j++)
|
||||
{
|
||||
|
||||
fscanf(pFile, "%f", &v);
|
||||
@@ -228,17 +262,23 @@ private:
|
||||
uint32_t m_numHidden;
|
||||
noise_gen_t m_noise;
|
||||
|
||||
void alloc()
|
||||
void alloc(uint32_t numVisible, uint32_t numHidden)
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
if (m_ppW)
|
||||
{
|
||||
if ((numVisible == m_numVisible) && (numHidden == m_numHidden))
|
||||
{
|
||||
return;
|
||||
}
|
||||
free();
|
||||
alloc();
|
||||
alloc(numVisible, numHidden);
|
||||
}
|
||||
else
|
||||
{
|
||||
m_numVisible = numVisible;
|
||||
m_numHidden = numHidden;
|
||||
m_ppW = new double*[m_numHidden];
|
||||
for (i=0; i < m_numHidden; i++)
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user