[RBM]
- removed classes HiddenLayer.hpp and VisibleLayer.hpp - fixed warnings - RbmComponent inherits Rbm - improved Rbm::train - use gaussion weight initialization git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@297 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+116
-89
@@ -8,10 +8,7 @@
|
||||
#ifndef RBM_HPP_
|
||||
#define RBM_HPP_
|
||||
|
||||
#include "VisibleLayer.hpp"
|
||||
#include "HiddenLayer.hpp"
|
||||
#include "Weights.hpp"
|
||||
#include "LayerArray.hpp"
|
||||
#include <cmath>
|
||||
#include <Eigen/Dense>
|
||||
|
||||
@@ -20,20 +17,7 @@ using namespace Eigen;
|
||||
void mylog(const char* format, ...);
|
||||
#define printf mylog
|
||||
|
||||
#define EPSILON_SIGMA 0.05
|
||||
|
||||
class Rbm;
|
||||
|
||||
class RbmListener
|
||||
{
|
||||
public:
|
||||
RbmListener() {}
|
||||
virtual ~RbmListener()
|
||||
{
|
||||
}
|
||||
|
||||
virtual void onProgressChanged(const Rbm &obj) = 0;
|
||||
};
|
||||
#define EPSILON_SIGMA 0.001
|
||||
|
||||
class Rbm
|
||||
{
|
||||
@@ -43,15 +27,16 @@ public:
|
||||
Params()
|
||||
: m_constantSigma(1.0)
|
||||
, m_sigmaDecay(1.0)
|
||||
, m_weightDecay(0.0)
|
||||
, m_weightDecay(0.01)
|
||||
, m_lambda(1.0)
|
||||
, m_sparsity(0.05)
|
||||
, m_muWeights(0.01)
|
||||
, m_muWeights(0.1)
|
||||
, m_muSparsity(0.01)
|
||||
, m_momentum(0.5)
|
||||
, m_useVisibleGaussian(false)
|
||||
, m_doRaoBlackwell(false)
|
||||
, m_useProbsForHiddenReconstruction(false)
|
||||
, m_doRaoBlackwell(true)
|
||||
, m_doSampleVisible(false)
|
||||
, m_doSampleBatch(false)
|
||||
, m_doSparse(false)
|
||||
, m_doNormalizeData(false)
|
||||
, m_doLearnVariance(false)
|
||||
@@ -69,18 +54,18 @@ public:
|
||||
double m_momentum;
|
||||
bool m_useVisibleGaussian;
|
||||
bool m_doRaoBlackwell;
|
||||
bool m_useProbsForHiddenReconstruction;
|
||||
bool m_doSampleVisible;
|
||||
bool m_doSampleBatch;
|
||||
bool m_doSparse;
|
||||
bool m_doNormalizeData;
|
||||
bool m_doLearnVariance;
|
||||
uint32_t m_numGibbs;
|
||||
size_t m_numGibbs;
|
||||
};
|
||||
|
||||
Rbm(Weights &weights, const MatrixXd &batch, RbmListener *pListener = nullptr)
|
||||
Rbm(Weights &weights, const MatrixXd &batch)
|
||||
: m_w(weights)
|
||||
, m_batch(batch)
|
||||
, m_variableSigma(weights.getNumVisible())
|
||||
, m_pListener(pListener)
|
||||
, m_progress(0)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
@@ -193,6 +178,16 @@ public:
|
||||
src.array() *= k.array();
|
||||
}
|
||||
|
||||
void sampleGaussian(MatrixXd &dst, MatrixXd const &src, const MatrixXd &sigma)
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < src.array().size(); i++)
|
||||
{
|
||||
dst.array()(i) = sigma(i)*Noise_Gaussian(&m_noise) + src.array()(i);
|
||||
}
|
||||
}
|
||||
|
||||
void sampleGaussian(MatrixXd &src, const MatrixXd &sigma)
|
||||
{
|
||||
uint32_t i;
|
||||
@@ -258,24 +253,27 @@ public:
|
||||
size_t batchSize = m_batch.rows();
|
||||
|
||||
double dProgress = 1.0/numEpochs;
|
||||
double kTrain = 1.0/batchSize;
|
||||
|
||||
double mu_w = m_params.m_muWeights/batchSize;
|
||||
double mu_biasV = m_params.m_muWeights/batchSize;
|
||||
double mu_biasH = m_params.m_muWeights/batchSize;
|
||||
|
||||
|
||||
m_v.resize(batchSize, m_w.getNumVisible());
|
||||
MatrixXd h(batchSize, m_w.getNumHidden());
|
||||
|
||||
MatrixXd sumBiasV(1, m_w.getNumVisible());
|
||||
MatrixXd sumBiasH(1, m_w.getNumHidden());
|
||||
MatrixXd sumWeights(m_w.getNumVisible(), m_w.getNumHidden());
|
||||
|
||||
MatrixXd deltaBiasV(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||
MatrixXd deltaBiasH(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||
MatrixXd deltaWeights(MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden()));
|
||||
MatrixXd dBiasV_curr(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||
MatrixXd dBiasH_curr(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||
MatrixXd dW_curr(MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden()));
|
||||
MatrixXd dBiasV(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||
MatrixXd dBiasH(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||
MatrixXd dW(MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden()));
|
||||
|
||||
MatrixXd diffErr(batchSize, m_w.getNumVisible());
|
||||
|
||||
MatrixXd batch = m_batch;
|
||||
|
||||
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
||||
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
||||
|
||||
if (m_params.m_doNormalizeData)
|
||||
{
|
||||
RowVectorXd mean = calcMean(batch);
|
||||
@@ -289,93 +287,97 @@ public:
|
||||
m_progress = 0;
|
||||
for (epoch=0; epoch < numEpochs; epoch++)
|
||||
{
|
||||
if (m_pListener)
|
||||
{
|
||||
m_pListener->onProgressChanged(*this);
|
||||
}
|
||||
onProgressChanged();
|
||||
|
||||
if (!m_params.m_useProbsForHiddenReconstruction)
|
||||
// When the hidden units are being driven by data, always use stochastic binary states
|
||||
if (m_params.m_doSampleBatch)
|
||||
{
|
||||
sample(batch, m_batch);
|
||||
}
|
||||
sample(batch_sampled, batch);
|
||||
|
||||
// Create hidden layer base on training data
|
||||
toHiddenBatch(h, batch);
|
||||
// Create hidden layer base on sampled training data
|
||||
toHiddenBatch(h, batch_sampled);
|
||||
}
|
||||
else
|
||||
{
|
||||
// Create hidden layer base on training data
|
||||
toHiddenBatch(h, batch);
|
||||
}
|
||||
// Sample hidden
|
||||
if (!m_params.m_doRaoBlackwell)
|
||||
{
|
||||
sample(h);
|
||||
}
|
||||
|
||||
// Update weights (positive phase)
|
||||
sumBiasV = batch.colwise().sum();
|
||||
dBiasV_curr = batch.colwise().sum();
|
||||
if (!m_params.m_doSparse)
|
||||
{
|
||||
sumBiasH = h.colwise().sum();
|
||||
dBiasH_curr = h.colwise().sum();
|
||||
}
|
||||
sumWeights = batch.transpose() * h;
|
||||
dW_curr = batch.transpose() * h;
|
||||
|
||||
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
|
||||
{
|
||||
sample(h);
|
||||
|
||||
// Create visible reconstruction (a fantasy...) given h
|
||||
toVisibleBatch(m_v, h);
|
||||
if (m_params.m_useVisibleGaussian)
|
||||
{
|
||||
if (!m_params.m_useProbsForHiddenReconstruction)
|
||||
{
|
||||
sampleGaussian(m_v, m_variableSigma.replicate(batchSize, 1));
|
||||
}
|
||||
sampleGaussian(v_sampled, m_v, m_variableSigma.replicate(batchSize, 1));
|
||||
toHiddenBatch(h, v_sampled);
|
||||
}
|
||||
else
|
||||
{
|
||||
probsLogistic(m_v, m_variableSigma.replicate(batchSize, 1));
|
||||
if (!m_params.m_useProbsForHiddenReconstruction)
|
||||
if (m_params.m_doSampleVisible)
|
||||
{
|
||||
sample(m_v);
|
||||
sample(v_sampled, m_v);
|
||||
// Create hidden representation given sampled v
|
||||
toHiddenBatch(h, v_sampled);
|
||||
}
|
||||
else
|
||||
{
|
||||
// Create hidden representation given v
|
||||
toHiddenBatch(h, m_v);
|
||||
}
|
||||
}
|
||||
|
||||
// Create hidden representation given v
|
||||
toHiddenBatch(h, m_v);
|
||||
if (!m_params.m_doRaoBlackwell)
|
||||
{
|
||||
sample(h);
|
||||
}
|
||||
}
|
||||
|
||||
if (!m_params.m_doRaoBlackwell)
|
||||
{
|
||||
sample(h);
|
||||
}
|
||||
// Update weights (negative phase)
|
||||
sumBiasV -= m_v.colwise().sum();
|
||||
dBiasV_curr -= m_v.colwise().sum();
|
||||
if (!m_params.m_doSparse)
|
||||
{
|
||||
sumBiasH -= h.colwise().sum();
|
||||
dBiasH_curr -= h.colwise().sum();
|
||||
}
|
||||
sumWeights -= m_v.transpose() * h;
|
||||
|
||||
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_params.m_momentum*deltaBiasV + m_params.m_muWeights*kTrain*sumBiasV;
|
||||
m_w.visibleBias() += deltaBiasV;
|
||||
dW_curr -= m_v.transpose() * h;
|
||||
|
||||
m_w.visibleBias() += mu_biasV*(m_params.m_momentum*dBiasV + (1-m_params.m_momentum)*dBiasV_curr);
|
||||
dBiasV = dBiasV_curr;
|
||||
|
||||
m_w.weights() += mu_w*(m_params.m_momentum*dW + (1-m_params.m_momentum)*dW_curr - m_params.m_weightDecay*m_w.weights());
|
||||
dW = dW_curr;
|
||||
|
||||
if (m_params.m_doSparse)
|
||||
{
|
||||
// Create hidden representation given v
|
||||
toHiddenBatch(h, m_v);
|
||||
toHiddenBatch(h, batch);
|
||||
|
||||
sumBiasH.fill(m_params.m_sparsity);
|
||||
sumBiasH -= h.colwise().mean();
|
||||
|
||||
deltaBiasH = m_params.m_momentum*deltaBiasH + m_params.m_muSparsity*sumBiasH;
|
||||
dBiasH_curr = m_params.m_sparsity * MatrixXd::Ones(dBiasH.rows(), dBiasH.cols()) - h.colwise().mean();
|
||||
|
||||
m_w.hiddenBias() += m_params.m_muSparsity*(m_params.m_momentum*dBiasH + (1-m_params.m_momentum)*dBiasH_curr);
|
||||
dBiasH = dBiasH_curr;
|
||||
|
||||
// cout << "Mean(" << m_sparsity << ") = " << (double)sumBiasH.array().mean() << endl;
|
||||
// cout << sumBiasH << endl;
|
||||
}
|
||||
else
|
||||
{
|
||||
deltaBiasH = m_params.m_momentum*deltaBiasH + m_params.m_muWeights*kTrain*sumBiasH;
|
||||
m_w.hiddenBias() += mu_biasH*(m_params.m_momentum*dBiasH + (1-m_params.m_momentum)*dBiasH_curr);
|
||||
dBiasH = dBiasH_curr;
|
||||
}
|
||||
m_w.hiddenBias() += deltaBiasH;
|
||||
|
||||
if (m_variableSigma[0] > sigmaMin)
|
||||
{
|
||||
@@ -383,12 +385,8 @@ public:
|
||||
}
|
||||
|
||||
m_progress += dProgress;
|
||||
if (m_pListener)
|
||||
{
|
||||
m_pListener->onProgressChanged(*this);
|
||||
}
|
||||
|
||||
diffErr = batch - m_v;
|
||||
diffErr = m_batch - m_v;
|
||||
diffErr.array() *= diffErr.array();
|
||||
double err = diffErr.colwise().sum().sum();
|
||||
|
||||
@@ -398,7 +396,7 @@ public:
|
||||
} // Number of epochs
|
||||
|
||||
updateHiddenBatch();
|
||||
|
||||
onProgressChanged();
|
||||
}
|
||||
|
||||
double getProgress() const
|
||||
@@ -432,7 +430,7 @@ public:
|
||||
|
||||
if (m_params.m_useVisibleGaussian)
|
||||
{
|
||||
// probsGaussian(v, m_sigmas);
|
||||
probsGaussian(v, m_variableSigma);
|
||||
}
|
||||
else
|
||||
{
|
||||
@@ -444,6 +442,7 @@ public:
|
||||
{
|
||||
m_params.m_constantSigma = value;
|
||||
m_variableSigma.fill(m_params.m_constantSigma);
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
RowVectorXd& getVariableSigma()
|
||||
@@ -454,46 +453,61 @@ public:
|
||||
void setSigmaDecay(double value)
|
||||
{
|
||||
m_params.m_sigmaDecay = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setWeightDecay(double value)
|
||||
{
|
||||
m_params.m_weightDecay = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setLambda(double value)
|
||||
{
|
||||
m_params.m_lambda = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setSparsity(double value)
|
||||
{
|
||||
m_params.m_sparsity = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setUseVisibleGaussian(bool flag)
|
||||
{
|
||||
m_params.m_useVisibleGaussian = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setDoRaoBlackwell(bool flag)
|
||||
{
|
||||
m_params.m_doRaoBlackwell = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setUseProbsForHiddenReconstruction(bool flag)
|
||||
void setDoSampleVisible(bool flag)
|
||||
{
|
||||
m_params.m_useProbsForHiddenReconstruction = flag;
|
||||
m_params.m_doSampleVisible = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setDoSampleBatch(bool flag)
|
||||
{
|
||||
m_params.m_doSampleBatch = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setDoSparse(bool flag)
|
||||
{
|
||||
m_params.m_doSparse = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setNormalizeData(bool flag)
|
||||
{
|
||||
m_params.m_doNormalizeData = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setDoLearnVariance(bool flag)
|
||||
@@ -507,26 +521,31 @@ public:
|
||||
{
|
||||
m_variableSigma.fill(m_params.m_constantSigma);
|
||||
}
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setNumGibbs(uint32_t value)
|
||||
void setNumGibbs(size_t value)
|
||||
{
|
||||
m_params.m_numGibbs = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setMuWeights(double value)
|
||||
{
|
||||
m_params.m_muWeights = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setMuSparsity(double value)
|
||||
{
|
||||
m_params.m_muSparsity = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void setMomentum(double value)
|
||||
{
|
||||
m_params.m_momentum = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
MatrixXd const& getHiddenBatch()
|
||||
@@ -561,7 +580,6 @@ private:
|
||||
MatrixXd m_v;
|
||||
MatrixXd m_h;
|
||||
RowVectorXd m_variableSigma;
|
||||
RbmListener *m_pListener;
|
||||
noise_gen_t m_noise;
|
||||
double m_progress;
|
||||
Params m_params;
|
||||
@@ -581,6 +599,15 @@ private:
|
||||
v = h * m_w.weights().transpose();
|
||||
v += m_w.visibleBias().replicate(m_batch.rows(), 1);
|
||||
}
|
||||
|
||||
protected:
|
||||
virtual void onProgressChanged()
|
||||
{
|
||||
}
|
||||
|
||||
virtual void onParamsChanged()
|
||||
{
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user