- added DBN stack
- concentrated RBM params into structure


git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@295 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-06-17 22:14:55 +00:00
parent b10d06d96c
commit 46bffda38f
5 changed files with 319 additions and 198 deletions
+125 -95
View File
@@ -38,27 +38,50 @@ public:
class Rbm
{
public:
struct Params
{
Params()
: m_constantSigma(1.0)
, m_sigmaDecay(1.0)
, m_weightDecay(0.0)
, m_lambda(1.0)
, m_sparsity(0.05)
, m_muWeights(0.01)
, m_muSparsity(0.01)
, m_momentum(0.5)
, m_useVisibleGaussian(false)
, m_doRaoBlackwell(false)
, m_useProbsForHiddenReconstruction(false)
, m_doSparse(false)
, m_doNormalizeData(false)
, m_doLearnVariance(false)
, m_numGibbs(1)
{
}
double m_constantSigma;
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_doRaoBlackwell;
bool m_useProbsForHiddenReconstruction;
bool m_doSparse;
bool m_doNormalizeData;
bool m_doLearnVariance;
uint32_t m_numGibbs;
};
Rbm(Weights &weights, const MatrixXd &batch, RbmListener *pListener = nullptr)
: m_w(weights)
, m_batch(batch)
, m_variableSigma(weights.getNumVisible())
, m_constantSigma(1.0)
, m_pListener(pListener)
, m_progress(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_useVisibleGaussian(false)
, m_doRaoBlackwell(false)
, m_useProbsForHiddenReconstruction(false)
, m_doSparse(false)
, m_doNormalizeData(false)
, m_doLearnVariance(false)
, m_numGibbs(1)
{
Noise_Init(&m_noise, 0x32727155);
@@ -72,6 +95,8 @@ public:
cout << b << endl;
#endif
m_variableSigma.fill(m_params.m_constantSigma);
updateHiddenBatch();
}
~Rbm()
@@ -79,13 +104,18 @@ public:
Noise_Free(&m_noise);
}
void sample(MatrixXd &src)
void sample(MatrixXd &srcDst)
{
sample(srcDst, srcDst);
}
void sample(MatrixXd &dst, MatrixXd const &src)
{
uint32_t i;
for (i=0; i < src.array().size(); i++)
{
src.array()(i) = (double)(src.array()(i) > Noise_Uniform(&m_noise));
dst.array()(i) = (double)(src.array()(i) > Noise_Uniform(&m_noise));
}
}
@@ -232,7 +262,7 @@ public:
m_v.resize(batchSize, m_w.getNumVisible());
m_h.resize(batchSize, m_w.getNumHidden());
MatrixXd h(batchSize, m_w.getNumHidden());
MatrixXd sumBiasV(1, m_w.getNumVisible());
MatrixXd sumBiasH(1, m_w.getNumHidden());
@@ -247,8 +277,8 @@ public:
m_progress = 0;
MatrixXd batch = m_batch;
if (m_doNormalizeData)
if (m_params.m_doNormalizeData)
{
RowVectorXd mean = calcMean(batch);
for (i=0; i < batchSize; i++)
@@ -257,38 +287,41 @@ public:
batch.row(i) = normalizeData(x, mean, m_variableSigma);
}
}
if (!m_params.m_useProbsForHiddenReconstruction)
{
sample(batch);
}
for (epoch=0; epoch < numEpochs; epoch++)
{
double err;
// Create hidden layer base on training data
toHiddenBatch(m_h, batch);
probsLogistic(m_h);
if (!m_doRaoBlackwell)
toHiddenBatch(h, batch);
if (!m_params.m_doRaoBlackwell)
{
sample(m_h);
sample(h);
}
// Update weights (positive phase)
sumBiasV = batch.colwise().sum();
if (!m_doSparse)
if (!m_params.m_doSparse)
{
sumBiasH = m_h.colwise().sum();
sumBiasH = h.colwise().sum();
}
sumWeights = batch.transpose() * m_h;
sumWeights = batch.transpose() * h;
diffErr = batch;
for (gibbs=0; gibbs < m_numGibbs; gibbs++)
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
{
sample(m_h);
sample(h);
// Create visible reconstruction (a fantasy...) given h
toVisibleBatch(m_v, m_h);
if (m_useVisibleGaussian)
toVisibleBatch(m_v, h);
if (m_params.m_useVisibleGaussian)
{
if (!m_useProbsForHiddenReconstruction)
if (!m_params.m_useProbsForHiddenReconstruction)
{
sampleGaussian(m_v, m_variableSigma.replicate(batchSize, 1));
}
@@ -296,59 +329,57 @@ public:
else
{
probsLogistic(m_v, m_variableSigma.replicate(batchSize, 1));
if (!m_useProbsForHiddenReconstruction)
if (!m_params.m_useProbsForHiddenReconstruction)
{
sample(m_v);
}
}
// Create hidden representation given v
toHiddenBatch(m_h, m_v);
probsLogistic(m_h);
toHiddenBatch(h, m_v);
}
if (!m_doRaoBlackwell)
if (!m_params.m_doRaoBlackwell)
{
sample(m_h);
sample(h);
}
// Update weights (negative phase)
sumBiasV -= m_v.colwise().sum();
if (!m_doSparse)
if (!m_params.m_doSparse)
{
sumBiasH -= m_h.colwise().sum();
sumBiasH -= h.colwise().sum();
}
sumWeights -= m_v.transpose() * m_h;
sumWeights -= m_v.transpose() * h;
diffErr -= m_v;
deltaWeights = m_momentum*deltaWeights + m_muWeights*(kTrain*sumWeights - 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;
deltaBiasV = m_momentum*deltaBiasV + m_muWeights*kTrain*sumBiasV;
deltaBiasV = m_params.m_momentum*deltaBiasV + m_params.m_muWeights*kTrain*sumBiasV;
m_w.visibleBias() += deltaBiasV;
if (m_doSparse)
if (m_params.m_doSparse)
{
// Create hidden representation given v
toHiddenBatch(m_h, m_v);
probsLogistic(m_h);
toHiddenBatch(h, m_v);
sumBiasH.fill(m_sparsity);
sumBiasH -= m_h.colwise().mean();
sumBiasH.fill(m_params.m_sparsity);
sumBiasH -= h.colwise().mean();
deltaBiasH = m_momentum*deltaBiasH + m_muSparsity*sumBiasH;
deltaBiasH = m_params.m_momentum*deltaBiasH + m_params.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_params.m_momentum*deltaBiasH + m_params.m_muWeights*kTrain*sumBiasH;
}
m_w.hiddenBias() += deltaBiasH;
if (m_variableSigma[0] > sigmaMin)
{
m_variableSigma.array() *= m_sigmaDecay;
m_variableSigma.array() *= m_params.m_sigmaDecay;
}
m_progress += dProgress;
@@ -364,6 +395,9 @@ public:
cout << err << endl;
} // Number of epochs
updateHiddenBatch();
}
double getProgress() const
@@ -395,7 +429,7 @@ public:
v = h * m_w.weights().transpose();
v += m_w.visibleBias();
if (m_useVisibleGaussian)
if (m_params.m_useVisibleGaussian)
{
// probsGaussian(v, m_sigmas);
}
@@ -407,13 +441,13 @@ public:
void setConstantSigma(double value)
{
m_constantSigma = value;
m_variableSigma.fill(m_constantSigma);
m_params.m_constantSigma = value;
m_variableSigma.fill(m_params.m_constantSigma);
}
double getConstantSigma()
{
return m_constantSigma;
return m_params.m_constantSigma;
}
RowVectorXd& getVariableSigma()
@@ -423,85 +457,80 @@ public:
void setSigmaDecay(double value)
{
m_sigmaDecay = value;
m_params.m_sigmaDecay = value;
}
void setWeightDecay(double value)
{
m_weightDecay = value;
m_params.m_weightDecay = value;
}
void setLambda(double value)
{
m_lambda = value;
m_params.m_lambda = value;
}
void setSparsity(double value)
{
m_sparsity = value;
m_params.m_sparsity = value;
}
void setUseVisibleGaussian(bool flag)
{
m_useVisibleGaussian = flag;
m_params.m_useVisibleGaussian = flag;
}
void setDoRaoBlackwell(bool flag)
{
m_doRaoBlackwell = flag;
m_params.m_doRaoBlackwell = flag;
}
void setUseProbsForHiddenReconstruction(bool flag)
{
m_useProbsForHiddenReconstruction = flag;
m_params.m_useProbsForHiddenReconstruction = flag;
}
void setDoSparse(bool flag)
{
m_doSparse = flag;
m_params.m_doSparse = flag;
}
void setNormalizeData(bool flag)
{
m_doNormalizeData = flag;
m_params.m_doNormalizeData = flag;
}
void setDoLearnVariance(bool flag)
{
m_doLearnVariance = flag;
if (m_doLearnVariance && m_batch.rows())
m_params.m_doLearnVariance = flag;
if (m_params.m_doLearnVariance && m_batch.rows())
{
m_variableSigma = calcSigma(m_batch);
}
else
{
m_variableSigma.fill(m_constantSigma);
m_variableSigma.fill(m_params.m_constantSigma);
}
}
void setNumGibbs(uint32_t value)
{
m_numGibbs = value;
}
uint32_t getNumGibbs()
{
return m_numGibbs;
m_params.m_numGibbs = value;
}
void setMuWeights(double value)
{
m_muWeights = value;
m_params.m_muWeights = value;
}
void setMuSparsity(double value)
{
m_muSparsity = value;
m_params.m_muSparsity = value;
}
void setMomentum(double value)
{
m_momentum = value;
m_params.m_momentum = value;
}
MatrixXd const& getHiddenBatch()
@@ -518,6 +547,17 @@ public:
{
return m_batch;
}
void updateHiddenBatch()
{
m_h.resize(m_batch.rows(), m_w.getNumHidden());
toHiddenBatch(m_h, m_batch);
}
Params const& params()
{
return m_params;
}
private:
Weights &m_w;
@@ -525,29 +565,19 @@ private:
MatrixXd m_v;
MatrixXd m_h;
RowVectorXd m_variableSigma;
double m_constantSigma;
RbmListener *m_pListener;
noise_gen_t m_noise;
double m_progress;
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_doRaoBlackwell;
bool m_useProbsForHiddenReconstruction;
bool m_doSparse;
bool m_doNormalizeData;
bool m_doLearnVariance;
uint32_t m_numGibbs;
Params m_params;
void toHiddenBatch(MatrixXd &h, MatrixXd const &v)
{
h = v * m_w.weights();
h += m_w.hiddenBias().replicate(m_batch.rows(), 1);
if (v.cols() == m_w.weights().rows())
{
h = v * m_w.weights();
h += m_w.hiddenBias().replicate(m_batch.rows(), 1);
probsLogistic(h);
}
}
void toVisibleBatch(MatrixXd &v, MatrixXd const &h)