- added getHiddenBatch() and getVisibleBatch()
- cleaned up

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@292 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-06-15 21:30:28 +00:00
parent 3a200cb46e
commit 4dccddd50f
3 changed files with 56 additions and 45 deletions
+46 -32
View File
@@ -38,7 +38,7 @@ public:
class Rbm
{
public:
Rbm(Weights &weights, const LayerArray &batch, RbmListener *pListener = nullptr)
Rbm(Weights &weights, const MatrixXd &batch, RbmListener *pListener = nullptr)
: m_w(weights)
, m_batch(batch)
, m_variableSigma(weights.getNumVisible())
@@ -224,13 +224,15 @@ public:
uint32_t t, i;
uint32_t epoch;
uint32_t gibbs;
size_t batchSize = m_batch.rows();
double dProgress = 1.0/numEpochs;
double kTrain = 1.0/m_batch.getSize();
double kTrain = 1.0/batchSize;
size_t batchSize = m_batch.getSize();
MatrixXd v(batchSize, m_w.getNumVisible());
MatrixXd h(batchSize, m_w.getNumHidden());
m_v.resize(batchSize, m_w.getNumVisible());
m_h.resize(batchSize, m_w.getNumHidden());
MatrixXd sumBiasV(1, m_w.getNumVisible());
MatrixXd sumBiasH(1, m_w.getNumHidden());
@@ -245,7 +247,7 @@ public:
m_progress = 0;
MatrixXd batch = m_batch.data();
MatrixXd batch = m_batch;
if (m_doNormalizeData)
{
@@ -262,62 +264,62 @@ public:
double err;
// Create hidden layer base on training data
toHiddenBatch(h, batch);
probsLogistic(h);
toHiddenBatch(m_h, batch);
probsLogistic(m_h);
if (!m_doRaoBlackwell)
{
sample(h);
sample(m_h);
}
// Update weights (positive phase)
sumBiasV = batch.colwise().sum();
if (!m_doSparse)
{
sumBiasH = h.colwise().sum();
sumBiasH = m_h.colwise().sum();
}
sumWeights = batch.transpose() * h;
sumWeights = batch.transpose() * m_h;
diffErr = batch;
for (gibbs=0; gibbs < m_numGibbs; gibbs++)
{
sample(h);
sample(m_h);
// Create visible reconstruction (a fantasy...) given h
toVisibleBatch(v, h);
toVisibleBatch(m_v, m_h);
if (m_useVisibleGaussian)
{
if (!m_useProbsForHiddenReconstruction)
{
sampleGaussian(v, m_variableSigma.replicate(batchSize, 1));
sampleGaussian(m_v, m_variableSigma.replicate(batchSize, 1));
}
}
else
{
probsLogistic(v, m_variableSigma.replicate(batchSize, 1));
probsLogistic(m_v, m_variableSigma.replicate(batchSize, 1));
if (!m_useProbsForHiddenReconstruction)
{
sample(v);
sample(m_v);
}
}
// Create hidden reconstruction given v
toHiddenBatch(h, v);
probsLogistic(h);
toHiddenBatch(m_h, m_v);
probsLogistic(m_h);
}
if (!m_doRaoBlackwell)
{
sample(h);
sample(m_h);
}
// Update weights (negative phase)
sumBiasV -= v.colwise().sum();
sumBiasV -= m_v.colwise().sum();
if (!m_doSparse)
{
sumBiasH -= h.colwise().sum();
sumBiasH -= m_h.colwise().sum();
}
sumWeights -= v.transpose() * h;
diffErr -= v;
sumWeights -= m_v.transpose() * m_h;
diffErr -= m_v;
deltaWeights = m_momentum*deltaWeights + m_muWeights*(kTrain*sumWeights - m_weightDecay*m_w.weights());
m_w.weights() += deltaWeights;
@@ -327,12 +329,12 @@ public:
if (m_doSparse)
{
h = v * m_w.weights();
h += m_w.hiddenBias().replicate(batchSize, 1);
probsLogistic(h);
m_h = m_v * m_w.weights();
m_h += m_w.hiddenBias().replicate(batchSize, 1);
probsLogistic(m_h);
sumBiasH.fill(m_sparsity);
sumBiasH -= h.colwise().mean();
sumBiasH -= m_h.colwise().mean();
deltaBiasH = m_momentum*deltaBiasH + m_muSparsity*sumBiasH;
@@ -468,9 +470,9 @@ public:
void setDoLearnVariance(bool flag)
{
m_doLearnVariance = flag;
if (m_doLearnVariance && m_batch.getSize())
if (m_doLearnVariance && m_batch.rows())
{
m_variableSigma = calcSigma(m_batch.data());
m_variableSigma = calcSigma(m_batch);
}
else
{
@@ -503,9 +505,21 @@ public:
m_momentum = value;
}
MatrixXd const& getHiddenBatch()
{
return m_h;
}
MatrixXd const& getVisibleBatch()
{
return m_v;
}
private:
Weights &m_w;
LayerArray const &m_batch;
MatrixXd const &m_batch;
MatrixXd m_v;
MatrixXd m_h;
RowVectorXd m_variableSigma;
double m_constantSigma;
RbmListener *m_pListener;
@@ -529,13 +543,13 @@ private:
void toHiddenBatch(MatrixXd &h, MatrixXd const &v)
{
h = v * m_w.weights();
h += m_w.hiddenBias().replicate(m_batch.getSize(), 1);
h += m_w.hiddenBias().replicate(m_batch.rows(), 1);
}
void toVisibleBatch(MatrixXd &v, MatrixXd const &h)
{
v = h * m_w.weights().transpose();
v += m_w.visibleBias().replicate(m_batch.getSize(), 1);
v += m_w.visibleBias().replicate(m_batch.rows(), 1);
}
};