- 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:
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user