- RBM: changed calcualation of pre-weight update data in RBM
- always use expectations - removed "Use Expectations Button" - removed Robbins-Monro - added sparsity learning rate - added momentum - added weight decay - added Slider as progress bar git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@30 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+94
-134
@@ -25,7 +25,9 @@ class RbmListener
|
||||
{
|
||||
public:
|
||||
RbmListener() {}
|
||||
virtual ~RbmListener() {}
|
||||
virtual ~RbmListener()
|
||||
{
|
||||
}
|
||||
|
||||
virtual void onEpochTrained(const Rbm &obj) = 0;
|
||||
};
|
||||
@@ -39,42 +41,29 @@ public:
|
||||
, m_progress(0)
|
||||
, m_sigma(1.0)
|
||||
, m_sigmaDecay(1.0)
|
||||
, m_weightDecay(0.0)
|
||||
, m_lambda(1.0)
|
||||
, m_sparsity(0)
|
||||
, m_muWeights(0.01)
|
||||
, m_muSparsity(0.01)
|
||||
, m_momentum(0.5)
|
||||
, m_doCancel(false)
|
||||
, m_useVisibleGaussian(false)
|
||||
, m_useExpectations(false)
|
||||
, m_doRaoBlackwell(false)
|
||||
, m_useProbsForHiddenReconstruction(false)
|
||||
, m_doRobbinsMonro(false)
|
||||
, m_doSparse(false)
|
||||
, m_numGibbs(1)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
}
|
||||
|
||||
~Rbm()
|
||||
{
|
||||
cancel();
|
||||
Noise_Free(&m_noise);
|
||||
}
|
||||
|
||||
void weightsUpdate(VisibleLayer &v, HiddenLayer &h, double mu)
|
||||
{
|
||||
MatrixXd &w = (MatrixXd&)m_w.weights();
|
||||
|
||||
// Update weights
|
||||
w += mu*(v.states() * h.states().transpose());
|
||||
}
|
||||
|
||||
void visibleBiasUpdate(VisibleLayer &v, double mu)
|
||||
{
|
||||
m_w.visibleBias().array() += mu*v.states().array();
|
||||
}
|
||||
|
||||
void hiddenBiasUpdate(HiddenLayer &h, double mu)
|
||||
{
|
||||
m_w.hiddenBias().array() += mu*h.states().array();
|
||||
}
|
||||
|
||||
void train(LayerArray<VisibleLayer> &vt, uint32_t numEpochs, double mu, uint32_t numGibbs = 1, double sigmaMin = 0.05)
|
||||
void train(LayerArray<VisibleLayer> &vt, uint32_t numEpochs, double sigmaMin = 0.05)
|
||||
{
|
||||
uint32_t t, i;
|
||||
uint32_t epoch;
|
||||
@@ -82,77 +71,72 @@ public:
|
||||
double sigma;
|
||||
VisibleLayer v(m_w.getNumVisible());
|
||||
HiddenLayer h(m_w.getNumHidden());
|
||||
HiddenLayer *pH;
|
||||
LayerArray<HiddenLayer> ht(vt.getSize(), m_w.getNumHidden());
|
||||
|
||||
VectorXd sumBiasV(m_w.getNumVisible());
|
||||
VectorXd deltaBiasV(m_w.getNumVisible());
|
||||
|
||||
VectorXd sumBiasH(m_w.getNumHidden());
|
||||
VectorXd deltaBiasH(m_w.getNumHidden());
|
||||
|
||||
MatrixXd sumWeights(m_w.getNumVisible(), m_w.getNumHidden());
|
||||
MatrixXd deltaWeights(m_w.getNumVisible(), m_w.getNumHidden());
|
||||
|
||||
sigma = m_sigma;
|
||||
Weights w = m_w;
|
||||
|
||||
double dProgress = 1.0/numEpochs;
|
||||
double kTrain = 1.0/vt.getSize();
|
||||
|
||||
m_progress = 0;
|
||||
|
||||
if (m_useExpectations)
|
||||
{
|
||||
mu /= vt.getSize();
|
||||
}
|
||||
|
||||
if (m_doRobbinsMonro)
|
||||
{
|
||||
for (i=0; i < vt.getSize(); i++)
|
||||
{
|
||||
// Create hidden layer base on training data
|
||||
ht[i].probsUpdateLogistic(vt[i], w, m_lambda, sigma);
|
||||
}
|
||||
}
|
||||
|
||||
deltaWeights.fill(0);
|
||||
deltaBiasV.fill(0);
|
||||
deltaBiasH.fill(0);
|
||||
m_doCancel = false;
|
||||
for (epoch=0; epoch < numEpochs; epoch++)
|
||||
{
|
||||
if (m_doCancel)
|
||||
{
|
||||
m_doCancel = false;
|
||||
break;
|
||||
}
|
||||
|
||||
sumWeights.fill(0);
|
||||
sumBiasV.fill(0);
|
||||
sumBiasH.fill(0);
|
||||
for (i=0; i < vt.getSize(); i++)
|
||||
{
|
||||
//t = (uint32_t)(0.5 + (vt.getSize()-1)*Noise_Uniform(&m_noise, 0.5));
|
||||
t = i;
|
||||
h.probsUpdateLogistic(vt[t], w, m_lambda, sigma);
|
||||
h.probsUpdateLogistic(vt[t], m_w, m_lambda, sigma);
|
||||
|
||||
// Create hidden layer base on training data
|
||||
if (m_doRobbinsMonro)
|
||||
{
|
||||
pH = &ht[t];
|
||||
}
|
||||
else
|
||||
{
|
||||
pH = &h;
|
||||
}
|
||||
|
||||
// Update weights (positive phase)
|
||||
if (m_doRaoBlackwell)
|
||||
{
|
||||
pH->states() = pH->probs();
|
||||
h.states() = h.probs();
|
||||
}
|
||||
else
|
||||
{
|
||||
pH->statesUpdateStochastic();
|
||||
}
|
||||
weightsUpdate(vt[t], h, +mu);
|
||||
visibleBiasUpdate(vt[t], +mu);
|
||||
if (!m_doSparse)
|
||||
{
|
||||
hiddenBiasUpdate(h, +mu);
|
||||
h.statesUpdateStochastic();
|
||||
}
|
||||
// Update weights (positive phase)
|
||||
sumWeights += vt[t].states() * h.states().transpose();
|
||||
sumBiasV += vt[t].states();
|
||||
sumBiasH += h.states();
|
||||
|
||||
for (gibbs=0; gibbs < numGibbs; gibbs++)
|
||||
for (gibbs=0; gibbs < m_numGibbs; gibbs++)
|
||||
{
|
||||
pH->statesUpdateStochastic();
|
||||
h.statesUpdateStochastic();
|
||||
|
||||
// Create visible reconstruction (a fantasy...)
|
||||
if (m_useProbsForHiddenReconstruction)
|
||||
{
|
||||
if (m_useVisibleGaussian)
|
||||
{
|
||||
v.probsUpdateGaussian(*pH, w, m_lambda, sigma);
|
||||
v.probsUpdateGaussian(h, m_w, m_lambda, sigma);
|
||||
}
|
||||
else
|
||||
{
|
||||
v.probsUpdateLogistic(*pH, w, m_lambda, sigma);
|
||||
v.probsUpdateLogistic(h, m_w, m_lambda, sigma);
|
||||
}
|
||||
v.states() = v.probs();
|
||||
}
|
||||
@@ -160,43 +144,38 @@ public:
|
||||
{
|
||||
if (m_useVisibleGaussian)
|
||||
{
|
||||
v.sampleGaussian(*pH, w, m_lambda, sigma);
|
||||
v.sampleGaussian(h, m_w, m_lambda, sigma);
|
||||
}
|
||||
else
|
||||
{
|
||||
v.probsUpdateLogistic(*pH, w, m_lambda, sigma);
|
||||
v.probsUpdateLogistic(h, m_w, m_lambda, sigma);
|
||||
v.statesUpdateStochastic();
|
||||
}
|
||||
}
|
||||
// Create hidden reconstruction
|
||||
pH->probsUpdateLogistic(v, w, m_lambda, sigma);
|
||||
h.probsUpdateLogistic(v, m_w, m_lambda, sigma);
|
||||
}
|
||||
|
||||
// Update weights (negative phase)
|
||||
if (m_doRaoBlackwell)
|
||||
{
|
||||
pH->states() = pH->probs();
|
||||
h.states() = h.probs();
|
||||
}
|
||||
else
|
||||
{
|
||||
pH->statesUpdateStochastic();
|
||||
}
|
||||
weightsUpdate(v, *pH, -mu);
|
||||
visibleBiasUpdate(v, -mu);
|
||||
if (!m_doSparse)
|
||||
{
|
||||
hiddenBiasUpdate(*pH, -mu);
|
||||
}
|
||||
if (!m_useExpectations)
|
||||
{
|
||||
w = m_w;
|
||||
h.statesUpdateStochastic();
|
||||
}
|
||||
sumWeights -= v.states() * h.states().transpose();
|
||||
sumBiasV -= v.states();
|
||||
sumBiasH -= h.states();
|
||||
|
||||
} // TrainingSize
|
||||
|
||||
if (m_useExpectations)
|
||||
{
|
||||
w = m_w;
|
||||
}
|
||||
deltaWeights = m_momentum*deltaWeights + m_muWeights*kTrain*sumWeights - m_weightDecay*m_w.weights();
|
||||
m_w.weights() += deltaWeights;
|
||||
|
||||
deltaBiasV = m_momentum*deltaBiasV + m_muWeights*kTrain*sumBiasV;
|
||||
m_w.visibleBias() += deltaBiasV;
|
||||
|
||||
if (m_doSparse)
|
||||
{
|
||||
@@ -206,17 +185,23 @@ public:
|
||||
|
||||
for (i=0; i < vt.getSize(); i++)
|
||||
{
|
||||
th.probsUpdateLogistic(vt[i], w, m_lambda, sigma);
|
||||
th.probsUpdateLogistic(vt[i], m_w, m_lambda, sigma);
|
||||
m += th.probs();
|
||||
}
|
||||
m /= i;
|
||||
th.states().array() = m.array() - m_sparsity;
|
||||
hiddenBiasUpdate(th, -mu);
|
||||
w = m_w;
|
||||
sumBiasH = m_sparsity - m.array();
|
||||
deltaBiasH = m_momentum*deltaBiasH + m_muSparsity*sumBiasH;
|
||||
|
||||
// cout << "Mean(" << m_sparsity << ") = " << (double)m.array().mean() << endl;
|
||||
// cout << m << endl;
|
||||
cout << "Mean(" << m_sparsity << ") = " << (double)m.array().mean() << endl;
|
||||
cout << m << endl;
|
||||
}
|
||||
else
|
||||
{
|
||||
deltaBiasH = m_momentum*deltaBiasH + m_muWeights*kTrain*sumBiasH;
|
||||
}
|
||||
|
||||
m_w.hiddenBias() += deltaBiasH;
|
||||
|
||||
if (sigma > sigmaMin)
|
||||
{
|
||||
sigma *= m_sigmaDecay;
|
||||
@@ -398,6 +383,11 @@ public:
|
||||
m_sigmaDecay = value;
|
||||
}
|
||||
|
||||
void setWeightDecay(double value)
|
||||
{
|
||||
m_weightDecay = value;
|
||||
}
|
||||
|
||||
void setLambda(double value)
|
||||
{
|
||||
m_lambda = value;
|
||||
@@ -413,11 +403,6 @@ public:
|
||||
m_useVisibleGaussian = flag;
|
||||
}
|
||||
|
||||
void setUseExpectations(bool flag)
|
||||
{
|
||||
m_useExpectations = flag;
|
||||
}
|
||||
|
||||
void setDoRaoBlackwell(bool flag)
|
||||
{
|
||||
m_doRaoBlackwell = flag;
|
||||
@@ -428,64 +413,35 @@ public:
|
||||
m_useProbsForHiddenReconstruction = flag;
|
||||
}
|
||||
|
||||
void setDoRobbinsMonro(bool flag)
|
||||
{
|
||||
m_doRobbinsMonro = flag;
|
||||
}
|
||||
|
||||
void setDoSparse(bool flag)
|
||||
{
|
||||
m_doSparse = flag;
|
||||
}
|
||||
|
||||
double getSparsity()
|
||||
void setNumGibbs(uint32_t value)
|
||||
{
|
||||
return m_sparsity;
|
||||
m_numGibbs = value;
|
||||
}
|
||||
|
||||
double getSigma()
|
||||
void setMuWeights(double value)
|
||||
{
|
||||
return m_sigma;
|
||||
m_muWeights = value;
|
||||
}
|
||||
|
||||
double getSigmaDecay()
|
||||
void setMuSparsity(double value)
|
||||
{
|
||||
return m_sigmaDecay;
|
||||
m_muSparsity = value;
|
||||
}
|
||||
|
||||
double getLambda()
|
||||
void setMomentum(double value)
|
||||
{
|
||||
return m_lambda;
|
||||
m_momentum = value;
|
||||
}
|
||||
|
||||
bool getUseVisibleGaussian()
|
||||
void cancel()
|
||||
{
|
||||
return m_useVisibleGaussian;
|
||||
}
|
||||
|
||||
bool getUseExpectations()
|
||||
{
|
||||
return m_useExpectations;
|
||||
}
|
||||
|
||||
bool getDoRaoBlackwell()
|
||||
{
|
||||
return m_doRaoBlackwell;
|
||||
}
|
||||
|
||||
bool getUseProbsForHiddenReconstruction()
|
||||
{
|
||||
return m_useProbsForHiddenReconstruction;
|
||||
}
|
||||
|
||||
bool getRobbinsMonro()
|
||||
{
|
||||
return m_doRobbinsMonro;
|
||||
}
|
||||
|
||||
bool getDoSparse()
|
||||
{
|
||||
return m_doSparse;
|
||||
m_doCancel = true;
|
||||
// while(m_doCancel);
|
||||
}
|
||||
|
||||
private:
|
||||
@@ -495,14 +451,18 @@ private:
|
||||
double m_progress;
|
||||
double m_sigma;
|
||||
double m_sigmaDecay;
|
||||
double m_weightDecay;
|
||||
double m_lambda;
|
||||
double m_sparsity;
|
||||
double m_muWeights;
|
||||
double m_muSparsity;
|
||||
double m_momentum;
|
||||
bool m_useVisibleGaussian;
|
||||
bool m_useExpectations;
|
||||
bool m_doRaoBlackwell;
|
||||
bool m_useProbsForHiddenReconstruction;
|
||||
bool m_doRobbinsMonro;
|
||||
bool m_doSparse;
|
||||
volatile bool m_doCancel;
|
||||
uint32_t m_numGibbs;
|
||||
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user