- DBN fixes
- batch sample inside training loop
- implemented RbmComponent stacking

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@296 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-06-18 17:57:32 +00:00
parent 46bffda38f
commit 1705287e37
5 changed files with 117 additions and 44 deletions
+23 -8
View File
@@ -508,7 +508,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
//[UserButtonCode_loadButton] -- add your button handler code here.. //[UserButtonCode_loadButton] -- add your button handler code here..
m_rbmSelect->setSelectedId(1, sendNotification); m_rbmSelect->setSelectedId(1, sendNotification);
create(getBaseDir()); create(getBaseDir());
//[/UserButtonCode_loadButton] //[/UserButtonCode_loadButton]
} }
else if (buttonThatWasClicked == saveButton) else if (buttonThatWasClicked == saveButton)
{ {
@@ -520,6 +520,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
{ {
//[UserButtonCode_loadTrainingButton] -- add your button handler code here.. //[UserButtonCode_loadTrainingButton] -- add your button handler code here..
loadTraining((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8()); loadTraining((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
m_pRbmComponent[0]->batchchanged();
//[/UserButtonCode_loadTrainingButton] //[/UserButtonCode_loadTrainingButton]
} }
else if (buttonThatWasClicked == saveTrainingButton) else if (buttonThatWasClicked == saveTrainingButton)
@@ -532,12 +533,14 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
{ {
//[UserButtonCode_clearTrainingButton] -- add your button handler code here.. //[UserButtonCode_clearTrainingButton] -- add your button handler code here..
clearTraining(); clearTraining();
m_pRbmComponent[0]->batchchanged();
//[/UserButtonCode_clearTrainingButton] //[/UserButtonCode_clearTrainingButton]
} }
else if (buttonThatWasClicked == removeTrainingButton) else if (buttonThatWasClicked == removeTrainingButton)
{ {
//[UserButtonCode_removeTrainingButton] -- add your button handler code here.. //[UserButtonCode_removeTrainingButton] -- add your button handler code here..
removeTrainingAt((int)patterSlider->getValue()); removeTrainingAt((int)patterSlider->getValue());
m_pRbmComponent[0]->batchchanged();
//[/UserButtonCode_removeTrainingButton] //[/UserButtonCode_removeTrainingButton]
} }
else if (buttonThatWasClicked == reconstructEquButton) else if (buttonThatWasClicked == reconstructEquButton)
@@ -798,14 +801,18 @@ void MainComponent::save ()
void MainComponent::create(juce::String const &projectName) void MainComponent::create(juce::String const &projectName)
{ {
size_t id = m_rbmSelect->getSelectedId()-1; size_t id = m_rbmSelect->getSelectedId()-1;
if ((m_pRbmComponent[id] == nullptr) and ((id+1) < DBN_SIZE)) bool shouldAddItem = true;
m_rbmSelect->addItem(String(id+1), id+2);
m_weights[id] = nullptr;
m_pRbmComponent[id] = nullptr;
if (id == 0) 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()) if (!projectName.isEmpty())
{ {
m_weights[id] = new Weights((String(projectName + String(".weights.dat"))).toUTF8()); m_weights[id] = new Weights((String(projectName + String(".weights.dat"))).toUTF8());
@@ -819,14 +826,22 @@ void MainComponent::create(juce::String const &projectName)
} }
else 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()); 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)); addAndMakeVisible(m_pRbmComponent[id] = new RbmComponent(*m_weights[id], m_pRbmComponent[id-1]->getHiddenBatch(), *this));
m_pRbmComponent[id-1]->registerRbm(m_pRbmComponent[id]);
} }
m_pRbmComponentCurr = m_pRbmComponent[id]; m_pRbmComponentCurr = m_pRbmComponent[id];
m_weightsCurr = m_weights[id]; m_weightsCurr = m_weights[id];
m_pRbmComponentCurr->setBounds (16, 140*id+16, 430, 130); m_pRbmComponentCurr->setBounds (16, 140*id+16, 430, 130);
m_pRbmComponentCurr->batchchanged(); m_pRbmComponentCurr->batchchanged();
m_pRbmComponentCurr->redrawReconstruction(); m_pRbmComponentCurr->redrawReconstruction();
if (((id+1) < DBN_SIZE) and shouldAddItem)
m_rbmSelect->addItem(String(id+1), id+2);
updateControls(); updateControls();
} }
@@ -854,7 +869,7 @@ void MainComponent::onChanged(const LayerArray &obj)
patterSlider->setRange(0, obj.getSize()-1, 1); patterSlider->setRange(0, obj.getSize()-1, 1);
} }
void MainComponent::onRbmEpochTrained(size_t progressPercent) void MainComponent::onProgressChanged(size_t progressPercent)
{ {
m_progressBarSlider->setValue(progressPercent); m_progressBarSlider->setValue(progressPercent);
} }
+1 -1
View File
@@ -54,7 +54,7 @@ public:
//============================================================================== //==============================================================================
//[UserMethods] -- You can add your own custom methods in this section. //[UserMethods] -- You can add your own custom methods in this section.
void onChanged(const LayerArray &obj) override; void onChanged(const LayerArray &obj) override;
void onRbmEpochTrained(size_t progressPercent) override; void onProgressChanged(size_t progressPercent) override;
//[/UserMethods] //[/UserMethods]
void paint (Graphics& g); void paint (Graphics& g);
+14 -18
View File
@@ -32,7 +32,7 @@ public:
{ {
} }
virtual void onEpochTrained(const Rbm &obj) = 0; virtual void onProgressChanged(const Rbm &obj) = 0;
}; };
class Rbm class Rbm
@@ -274,8 +274,6 @@ public:
MatrixXd diffErr(batchSize, m_w.getNumVisible()); MatrixXd diffErr(batchSize, m_w.getNumVisible());
m_progress = 0;
MatrixXd batch = m_batch; MatrixXd batch = m_batch;
if (m_params.m_doNormalizeData) if (m_params.m_doNormalizeData)
@@ -288,14 +286,18 @@ public:
} }
} }
if (!m_params.m_useProbsForHiddenReconstruction) m_progress = 0;
{
sample(batch);
}
for (epoch=0; epoch < numEpochs; epoch++) for (epoch=0; epoch < numEpochs; epoch++)
{ {
double err; if (m_pListener)
{
m_pListener->onProgressChanged(*this);
}
if (!m_params.m_useProbsForHiddenReconstruction)
{
sample(batch, m_batch);
}
// Create hidden layer base on training data // Create hidden layer base on training data
toHiddenBatch(h, batch); toHiddenBatch(h, batch);
@@ -311,7 +313,6 @@ public:
sumBiasH = h.colwise().sum(); sumBiasH = h.colwise().sum();
} }
sumWeights = batch.transpose() * h; sumWeights = batch.transpose() * h;
diffErr = batch;
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++) for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
{ {
@@ -350,7 +351,6 @@ public:
sumBiasH -= h.colwise().sum(); sumBiasH -= h.colwise().sum();
} }
sumWeights -= m_v.transpose() * h; sumWeights -= m_v.transpose() * h;
diffErr -= m_v;
deltaWeights = m_params.m_momentum*deltaWeights + m_params.m_muWeights*(kTrain*sumWeights - m_params.m_weightDecay*m_w.weights()); deltaWeights = m_params.m_momentum*deltaWeights + m_params.m_muWeights*(kTrain*sumWeights - m_params.m_weightDecay*m_w.weights());
m_w.weights() += deltaWeights; m_w.weights() += deltaWeights;
@@ -385,11 +385,12 @@ public:
m_progress += dProgress; m_progress += dProgress;
if (m_pListener) if (m_pListener)
{ {
m_pListener->onEpochTrained(*this); m_pListener->onProgressChanged(*this);
} }
diffErr = batch - m_v;
diffErr.array() *= diffErr.array(); diffErr.array() *= diffErr.array();
err = diffErr.colwise().sum().sum(); double err = diffErr.colwise().sum().sum();
cout << "err =" << endl; cout << "err =" << endl;
cout << err << endl; cout << err << endl;
@@ -445,11 +446,6 @@ public:
m_variableSigma.fill(m_params.m_constantSigma); m_variableSigma.fill(m_params.m_constantSigma);
} }
double getConstantSigma()
{
return m_params.m_constantSigma;
}
RowVectorXd& getVariableSigma() RowVectorXd& getVariableSigma()
{ {
return m_variableSigma; return m_variableSigma;
+15 -9
View File
@@ -204,20 +204,20 @@ void RbmComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails&
//[/UserCode_mouseWheelMove] //[/UserCode_mouseWheelMove]
} }
void RbmComponent::onEpochTrained(const Rbm &obj) void RbmComponent::onProgressChanged(const Rbm &obj)
{ {
// mylog("Training %f %%\n", obj.getProgress()*100); // mylog("Training %f %%\n", obj.getProgress()*100);
upPass(DrawHidden->getData());
redrawReconstruction(); redrawReconstruction();
redrawWeights(); redrawWeights();
m_listener.onRbmEpochTrained((size_t)(100*obj.getProgress())); m_listener.onProgressChanged((size_t)(100*obj.getProgress() + 0.5));
} }
void RbmComponent::onDraw(DrawComponent &obj) void RbmComponent::onDraw(DrawComponent &obj)
{ {
if (&obj == DrawHidden) if (&obj == DrawHidden)
{ {
m_pRbm->toVisible(DrawReconstruction->getData(), obj.getData()); upPass(obj.getData());
DrawReconstruction->DrawData();
} }
if (&obj == DrawTraining) if (&obj == DrawTraining)
{ {
@@ -229,15 +229,16 @@ void RbmComponent::redrawReconstruction()
{ {
uint32_t i; uint32_t i;
m_pRbm->toHidden(DrawHidden->getData(), DrawTraining->getData()); reconstructVisible(DrawReconstruction->getData(), DrawTraining->getData());
m_pRbm->toVisible(DrawReconstruction->getData(), DrawHidden->getData()); // toHidden(DrawHidden->getData(), DrawTraining->getData());
// toVisible(DrawReconstruction->getData(), DrawHidden->getData());
for (i=0; i < m_pRbm->params().m_numGibbs-1; i++) for (i=0; i < m_pRbm->params().m_numGibbs-1; i++)
{ {
m_pRbm->toHidden(DrawHidden->getData(), DrawReconstruction->getData()); reconstructVisible(DrawReconstruction->getData(), DrawReconstruction->getData());
m_pRbm->toVisible(DrawReconstruction->getData(), DrawHidden->getData()); // toHidden(DrawHidden->getData(), DrawReconstruction->getData());
// toVisible(DrawReconstruction->getData(), DrawHidden->getData());
} }
DrawHidden->DrawData();
DrawReconstruction->DrawData(); DrawReconstruction->DrawData();
} }
@@ -337,6 +338,10 @@ void RbmComponent::setMomentum(double value)
void RbmComponent::batchchanged() void RbmComponent::batchchanged()
{ {
m_pRbm->updateHiddenBatch(); m_pRbm->updateHiddenBatch();
if (lower)
{
lower->batchchanged();
}
redrawVariances(); redrawVariances();
} }
@@ -360,6 +365,7 @@ void RbmComponent::setTrainingIndex(size_t index)
DrawTraining->DrawData(); DrawTraining->DrawData();
redrawReconstruction(); redrawReconstruction();
upPass(DrawHidden->getData());
} }
} }
+63 -7
View File
@@ -27,13 +27,30 @@
//[/Headers] //[/Headers]
class IRbm class IRbmComponent
{ {
public: public:
IRbm() {} IRbmComponent *upper;
virtual ~IRbm() IRbmComponent *lower;
IRbmComponent()
: upper(nullptr)
, lower(nullptr)
{}
virtual ~IRbmComponent()
{ {
} }
virtual void batchchanged() = 0;
virtual void toHidden(RowVectorXd &h, RowVectorXd const &v) = 0;
virtual void toVisible(RowVectorXd &v, RowVectorXd const &h) = 0;
virtual void downPass(RowVectorXd &h) = 0;
virtual void upPass(RowVectorXd const &v) = 0;
void registerRbm(IRbmComponent *pObj)
{
lower = pObj;
pObj->upper = this;
}
}; };
class RbmComponentListener class RbmComponentListener
@@ -43,7 +60,7 @@ public:
virtual ~RbmComponentListener() virtual ~RbmComponentListener()
{ {
} }
virtual void onRbmEpochTrained(size_t progressPercent) = 0; virtual void onProgressChanged(size_t progressPercent) = 0;
}; };
@@ -57,7 +74,8 @@ public:
*/ */
class RbmComponent : public Component, class RbmComponent : public Component,
public RbmListener, public RbmListener,
public DrawListener public DrawListener,
public IRbmComponent
{ {
public: public:
//============================================================================== //==============================================================================
@@ -94,7 +112,7 @@ public:
void setMuWeights(double value); void setMuWeights(double value);
void setMuSparsity(double value); void setMuSparsity(double value);
void setMomentum(double value); void setMomentum(double value);
void batchchanged(); void batchchanged() override;
void setWeightsIndex(size_t index); void setWeightsIndex(size_t index);
size_t getWeightsIndex(); size_t getWeightsIndex();
@@ -117,6 +135,44 @@ public:
return m_pRbm->params(); return m_pRbm->params();
} }
void reconstructVisible(RowVectorXd& dst, const RowVectorXd& src)
{
toHidden(DrawHidden->getData(), src);
toVisible(dst, DrawHidden->getData());
DrawHidden->DrawData();
}
void toVisible(RowVectorXd &v, RowVectorXd const &h) override
{
m_pRbm->toVisible(v, h);
}
void toHidden(RowVectorXd &h, RowVectorXd const &v) override
{
m_pRbm->toHidden(h, v);
if (lower)
{
lower->downPass(h);
}
}
void downPass(RowVectorXd& h) override
{
toHidden(DrawHidden->getData(), h);
m_pRbm->toVisible(h, DrawHidden->getData());
DrawHidden->DrawData();
}
void upPass(RowVectorXd const &h) override
{
toVisible(DrawReconstruction->getData(), h);
DrawReconstruction->DrawData();
if (upper)
{
upper->upPass(DrawReconstruction->getData());
}
}
private: private:
//[UserVariables] -- You can add your own custom variables in this section. //[UserVariables] -- You can add your own custom variables in this section.
ScopedPointer<Rbm> m_pRbm; ScopedPointer<Rbm> m_pRbm;
@@ -129,7 +185,7 @@ private:
ScopedPointer<DrawComponent> DrawHidden; ScopedPointer<DrawComponent> DrawHidden;
size_t m_currWeightIndexToDraw; size_t m_currWeightIndexToDraw;
size_t m_currTrainingIndexToDraw; size_t m_currTrainingIndexToDraw;
void onEpochTrained(const Rbm &obj) override; void onProgressChanged(const Rbm &obj) override;
void onDraw(DrawComponent &obj) override; void onDraw(DrawComponent &obj) override;
//[/UserVariables] //[/UserVariables]