- train(): added reconstruction error metric

- train2(): 
  - added reconstruction error metric. 
  - added Gaussian units
  - added sparsity

- GUI: choose train() or train2()


git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@44 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2014-11-01 18:01:04 +00:00
parent 93fa8c64ce
commit 6e7b8bc151
3 changed files with 145 additions and 76 deletions
+24 -3
View File
@@ -280,6 +280,10 @@ MainComponent::MainComponent ()
weightInitLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000)); weightInitLabel->setColour (TextEditor::backgroundColourId, Colour (0x00000000));
weightInitLabel->addListener (this); weightInitLabel->addListener (this);
addAndMakeVisible (rbmTrainV2ToggleButton = new ToggleButton ("rbmTrainV2ToggleButton toggle button"));
rbmTrainV2ToggleButton->setButtonText (TRANS("Train Ver. 2"));
rbmTrainV2ToggleButton->addListener (this);
//[UserPreSize] //[UserPreSize]
m_vNumX = 16; m_vNumX = 16;
@@ -340,6 +344,7 @@ MainComponent::~MainComponent()
momentumLabel = nullptr; momentumLabel = nullptr;
sparsityLearningRateLabel = nullptr; sparsityLearningRateLabel = nullptr;
weightInitLabel = nullptr; weightInitLabel = nullptr;
rbmTrainV2ToggleButton = nullptr;
//[Destructor]. You can add your own custom destruction code here.. //[Destructor]. You can add your own custom destruction code here..
@@ -461,6 +466,7 @@ void MainComponent::resized()
momentumLabel->setBounds (208, 368, 72, 24); momentumLabel->setBounds (208, 368, 72, 24);
sparsityLearningRateLabel->setBounds (416, 320, 72, 24); sparsityLearningRateLabel->setBounds (416, 320, 72, 24);
weightInitLabel->setBounds (416, 368, 72, 24); weightInitLabel->setBounds (416, 368, 72, 24);
rbmTrainV2ToggleButton->setBounds (336, 200, 128, 24);
//[UserResized] Add your own custom resize handling here.. //[UserResized] Add your own custom resize handling here..
DrawTraining->setBounds (16, 16, 100, 100); DrawTraining->setBounds (16, 16, 100, 100);
DrawReconstruction->setBounds (110+16, 16, 100, 100); DrawReconstruction->setBounds (110+16, 16, 100, 100);
@@ -618,6 +624,11 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
m_pRbm->setDoSparse(buttonThatWasClicked->getToggleState()); m_pRbm->setDoSparse(buttonThatWasClicked->getToggleState());
//[/UserButtonCode_rbmDoSparseToggleButton] //[/UserButtonCode_rbmDoSparseToggleButton]
} }
else if (buttonThatWasClicked == rbmTrainV2ToggleButton)
{
//[UserButtonCode_rbmTrainV2ToggleButton] -- add your button handler code here..
//[/UserButtonCode_rbmTrainV2ToggleButton]
}
//[UserbuttonClicked_Post] //[UserbuttonClicked_Post]
//[/UserbuttonClicked_Post] //[/UserbuttonClicked_Post]
@@ -848,8 +859,8 @@ void MainComponent::redrawReconstruction()
{ {
DrawHidden->setData(m_pRbm->toHidden(DrawTraining->getData())); DrawHidden->setData(m_pRbm->toHidden(DrawTraining->getData()));
DrawReconstruction->setData(m_pRbm->toVisible(DrawHidden->getData())); DrawReconstruction->setData(m_pRbm->toVisible(DrawHidden->getData()));
double energy = m_pRbm->getEnergy(DrawTraining->getData(), DrawHidden->getData()); // double energy = m_pRbm->getEnergy(DrawTraining->getData(), DrawHidden->getData());
cout << "Energy(" << 0 <<") = " << energy << endl; // cout << "Energy(" << 0 <<") = " << energy << endl;
} }
@@ -862,8 +873,14 @@ void MainComponent::redrawWeights(int index)
void MainComponent::run() void MainComponent::run()
{ {
trainButton->setEnabled(false); trainButton->setEnabled(false);
if (rbmTrainV2ToggleButton->getToggleState())
{
m_pRbm->train2(m_layers, numEpochslabel->getText().getIntValue(), 100);
}
else
{
m_pRbm->train(m_layers, numEpochslabel->getText().getIntValue()); m_pRbm->train(m_layers, numEpochslabel->getText().getIntValue());
// m_pRbm->train2(m_layers, numEpochslabel->getText().getIntValue(), 100); }
trainButton->setEnabled(true); trainButton->setEnabled(true);
} }
@@ -1046,6 +1063,10 @@ BEGIN_JUCER_METADATA
edBkgCol="0" labelText="0.001" editableSingleClick="1" editableDoubleClick="1" edBkgCol="0" labelText="0.001" editableSingleClick="1" editableDoubleClick="1"
focusDiscardsChanges="0" fontname="Default font" fontsize="15" focusDiscardsChanges="0" fontname="Default font" fontsize="15"
bold="0" italic="0" justification="36"/> bold="0" italic="0" justification="36"/>
<TOGGLEBUTTON name="rbmTrainV2ToggleButton toggle button" id="92c647c1f8b110a2"
memberName="rbmTrainV2ToggleButton" virtualName="" explicitFocusOrder="0"
pos="336 200 128 24" buttonText="Train Ver. 2" connectedEdges="0"
needsCallback="1" radioGroupId="0" state="0"/>
</JUCER_COMPONENT> </JUCER_COMPONENT>
END_JUCER_METADATA END_JUCER_METADATA
+1
View File
@@ -129,6 +129,7 @@ private:
ScopedPointer<Label> momentumLabel; ScopedPointer<Label> momentumLabel;
ScopedPointer<Label> sparsityLearningRateLabel; ScopedPointer<Label> sparsityLearningRateLabel;
ScopedPointer<Label> weightInitLabel; ScopedPointer<Label> weightInitLabel;
ScopedPointer<ToggleButton> rbmTrainV2ToggleButton;
//============================================================================== //==============================================================================
+109 -62
View File
@@ -81,6 +81,8 @@ public:
MatrixXd sumWeights(m_w.getNumVisible(), m_w.getNumHidden()); MatrixXd sumWeights(m_w.getNumVisible(), m_w.getNumHidden());
MatrixXd deltaWeights(m_w.getNumVisible(), m_w.getNumHidden()); MatrixXd deltaWeights(m_w.getNumVisible(), m_w.getNumHidden());
MatrixXd diffErr(1, m_w.getNumVisible());
const LayerArray<VisibleLayer> &vt = batch; const LayerArray<VisibleLayer> &vt = batch;
sigma = m_sigma; sigma = m_sigma;
@@ -96,6 +98,8 @@ public:
m_doCancel = false; m_doCancel = false;
for (epoch=0; epoch < numEpochs; epoch++) for (epoch=0; epoch < numEpochs; epoch++)
{ {
double err = 0;
if (m_doCancel) if (m_doCancel)
{ {
m_doCancel = false; m_doCancel = false;
@@ -125,6 +129,8 @@ public:
sumBiasV += vt[t].states(); sumBiasV += vt[t].states();
sumBiasH += h.states(); sumBiasH += h.states();
diffErr = vt[t].states();
for (gibbs=0; gibbs < m_numGibbs; gibbs++) for (gibbs=0; gibbs < m_numGibbs; gibbs++)
{ {
h.statesUpdateStochastic(); h.statesUpdateStochastic();
@@ -170,6 +176,9 @@ public:
sumWeights -= v.states() * h.states().transpose(); sumWeights -= v.states() * h.states().transpose();
sumBiasV -= v.states(); sumBiasV -= v.states();
sumBiasH -= h.states(); sumBiasH -= h.states();
diffErr -= v.states();
diffErr.array() *= diffErr.array();
err += diffErr.sum();
} // TrainingSize } // TrainingSize
@@ -215,20 +224,40 @@ public:
m_pListener->onEpochTrained(*this); m_pListener->onEpochTrained(*this);
} }
cout << "err =" << endl;
cout << err << endl;
} // Number of epochs } // Number of epochs
} }
MatrixXd sample(const MatrixXd &src) void sample(MatrixXd &src)
{ {
uint32_t i; uint32_t i;
MatrixXd res(src);
for (i=0; i < src.array().size(); i++) for (i=0; i < src.array().size(); i++)
{ {
res.array()(i) = (double) src.array()(i) > Noise_Uniform(&m_noise); src.array()(i) = (double)(src.array()(i) > Noise_Uniform(&m_noise));
}
} }
return res; void probsLogistic(MatrixXd &src, double lambda, double sigma)
{
double var = sigma*sigma;
src.array() *= -lambda/var;
src.array() = src.array().exp();
src.array() += 1;
src.array() = 1.0/src.array();
}
void sampleGaussian(MatrixXd &src, double lambda, double sigma)
{
uint32_t i;
for (i=0; i < src.array().size(); i++)
{
src.array()(i) = sigma*Noise_Gaussian(&m_noise) + lambda*src.array()(i);
}
} }
void train2(const LayerArray<VisibleLayer> &vt, uint32_t numEpochs, uint32_t batchSize, double sigmaMin = 0.05) void train2(const LayerArray<VisibleLayer> &vt, uint32_t numEpochs, uint32_t batchSize, double sigmaMin = 0.05)
@@ -240,113 +269,125 @@ public:
double dProgress = 1.0/numEpochs; double dProgress = 1.0/numEpochs;
double kTrain = 1.0/vt.getSize(); double kTrain = 1.0/vt.getSize();
if (batchSize > vt.getSize()) // if (batchSize > vt.getSize())
batchSize = vt.getSize(); batchSize = vt.getSize();
MatrixXd vp(m_w.getNumVisible(), batchSize); MatrixXd v(batchSize, m_w.getNumVisible());
MatrixXd vs(m_w.getNumVisible(), batchSize); MatrixXd h(batchSize, m_w.getNumHidden());
MatrixXd hp(m_w.getNumHidden(), batchSize); MatrixXd batch(batchSize, m_w.getNumVisible());
MatrixXd hs(m_w.getNumHidden(), batchSize);
MatrixXd batch(m_w.getNumVisible(), batchSize);
VectorXd sumBiasV(m_w.getNumVisible()); MatrixXd sumBiasV(1, m_w.getNumVisible());
VectorXd sumBiasH(m_w.getNumHidden()); MatrixXd sumBiasH(1, m_w.getNumHidden());
MatrixXd sumWeights(m_w.getNumVisible(), m_w.getNumHidden()); MatrixXd sumWeights(m_w.getNumVisible(), m_w.getNumHidden());
VectorXd deltaBiasV(m_w.getNumVisible()); MatrixXd deltaBiasV(MatrixXd::Zero(1, m_w.getNumVisible()));
VectorXd deltaBiasH(m_w.getNumHidden()); MatrixXd deltaBiasH(MatrixXd::Zero(1, m_w.getNumHidden()));
MatrixXd deltaWeights(m_w.getNumVisible(), m_w.getNumHidden()); MatrixXd deltaWeights(MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden()));
MatrixXd diffErr(batchSize, m_w.getNumVisible());
m_progress = 0; m_progress = 0;
m_doCancel = false; m_doCancel = false;
for (i=0; i < batchSize; i++) for (i=0; i < batchSize; i++)
{ {
t = (uint32_t)(0.5 + (vt.getSize()-1)*Noise_Uniform(&m_noise)); // t = (uint32_t)(0.5 + (vt.getSize()-1)*Noise_Uniform(&m_noise));
batch.col(i) = vt[t].states(); batch.row(i) = vt[i].states();
} }
for (epoch=0; epoch < numEpochs; epoch++) for (epoch=0; epoch < numEpochs; epoch++)
{ {
double err;
v = batch;
if (m_doCancel) if (m_doCancel)
{ {
m_doCancel = false; m_doCancel = false;
break; break;
} }
sumWeights.fill(0);
sumBiasV.fill(0);
sumBiasH.fill(0);
for (i=0; i < batchSize; i++)
{
vs = batch.col(i);
// h.probsUpdateLogistic(vt[t], m_w, m_lambda, sigma);
hp = -vs.transpose() * m_w.weights();
hp.array() = hp.array().exp();
hp.array() += 1;
hp.array() = 1.0/hp.array();
// Create hidden layer base on training data // Create hidden layer base on training data
hs = sample(hp); h = v * m_w.weights();
h += m_w.hiddenBias().transpose().replicate(batchSize, 1);
probsLogistic(h, m_lambda, sigma);
if (!m_doRaoBlackwell)
{
sample(h);
}
// Update weights (positive phase) // Update weights (positive phase)
sumBiasV += vp.colwise().sum(); sumBiasV = v.colwise().sum();
if (m_doRaoBlackwell) if (!m_doSparse)
{ {
sumWeights += vs * hp.transpose(); sumBiasH = h.colwise().sum();
sumBiasH += hp.colwise().sum();
}
else
{
sumWeights += vs * hs.transpose();
sumBiasH += hs.colwise().sum();
} }
sumWeights = v.transpose() * h;
diffErr = v;
for (gibbs=0; gibbs < m_numGibbs; gibbs++) for (gibbs=0; gibbs < m_numGibbs; gibbs++)
{ {
sample(h);
// Create visible reconstruction (a fantasy...) // Create visible reconstruction (a fantasy...)
if (m_useProbsForHiddenReconstruction) v = h * m_w.weights().transpose();
v += m_w.visibleBias().transpose().replicate(batchSize, 1);
if (m_useVisibleGaussian)
{ {
vp = -hs * m_w.weights().transpose(); sampleGaussian(v, m_lambda, sigma);
vp.array() = vp.array().exp();
vp.array() += 1;
vp.array() = 1.0/vp.array();
} }
else else
{ {
vs = sample(vp); probsLogistic(v, m_lambda, sigma);
if (!m_useProbsForHiddenReconstruction)
{
sample(v);
}
} }
// Create hidden reconstruction // Create hidden reconstruction
hp = -vs.transpose() * m_w.weights(); h = v * m_w.weights();
hp.array() = hp.array().exp(); h += m_w.hiddenBias().transpose().replicate(batchSize, 1);
hp.array() += 1; probsLogistic(h, m_lambda, sigma);
hp.array() = 1.0/hp.array();
} }
if (!m_doRaoBlackwell)
{
sample(h);
}
// Update weights (negative phase) // Update weights (negative phase)
sumBiasV -= vp.colwise().sum(); sumBiasV -= v.colwise().sum();
if (m_doRaoBlackwell) if (!m_doSparse)
{ {
sumWeights -= vs * hp.transpose(); sumBiasH -= h.colwise().sum();
sumBiasH -= hp.colwise().sum();
}
else
{
sumWeights -= vs * hs.transpose();
sumBiasH -= hs.colwise().sum();
} }
sumWeights -= v.transpose() * h;
diffErr -= v;
} // TrainingSize deltaWeights = m_momentum*deltaWeights + m_muWeights*(kTrain*sumWeights - m_weightDecay*m_w.weights());
deltaWeights = m_momentum*deltaWeights + m_muWeights*kTrain*sumWeights - m_weightDecay*m_w.weights();
m_w.weights() += deltaWeights; m_w.weights() += deltaWeights;
deltaBiasV = m_momentum*deltaBiasV + m_muWeights*kTrain*sumBiasV; deltaBiasV = m_momentum*deltaBiasV + m_muWeights*kTrain*sumBiasV;
m_w.visibleBias() += deltaBiasV; m_w.visibleBias() += deltaBiasV;
if (m_doSparse)
{
h = v * m_w.weights();
h += m_w.hiddenBias().transpose().replicate(batchSize, 1);
probsLogistic(h, m_lambda, sigma);
sumBiasH.fill(m_sparsity);
sumBiasH -= h.colwise().mean();
deltaBiasH = m_momentum*deltaBiasH + m_muSparsity*sumBiasH;
// cout << "Mean(" << m_sparsity << ") = " << (double)sumBiasH.array().mean() << endl;
// cout << sumBiasH << endl;
}
else
{
deltaBiasH = m_momentum*deltaBiasH + m_muWeights*kTrain*sumBiasH; deltaBiasH = m_momentum*deltaBiasH + m_muWeights*kTrain*sumBiasH;
}
m_w.hiddenBias() += deltaBiasH; m_w.hiddenBias() += deltaBiasH;
if (sigma > sigmaMin) if (sigma > sigmaMin)
@@ -360,6 +401,12 @@ public:
m_pListener->onEpochTrained(*this); m_pListener->onEpochTrained(*this);
} }
diffErr.array() *= diffErr.array();
err = diffErr.colwise().sum().sum();
cout << "err =" << endl;
cout << err << endl;
} // Number of epochs } // Number of epochs
} }