[RBM]
- splitted bm into hpp and cpp git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@302 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+533
@@ -0,0 +1,533 @@
|
|||||||
|
/*
|
||||||
|
* To change this license header, choose License Headers in Project Properties.
|
||||||
|
* To change this template file, choose Tools | Templates
|
||||||
|
* and open the template in the editor.
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include "Rbm.hpp"
|
||||||
|
|
||||||
|
void mylog(const char* format, ...);
|
||||||
|
#define printf mylog
|
||||||
|
|
||||||
|
#define EPSILON_SIGMA 0.001
|
||||||
|
|
||||||
|
Rbm::Rbm(Weights &weights, const MatrixXd &batch)
|
||||||
|
: m_w(weights)
|
||||||
|
, m_batch(batch)
|
||||||
|
, m_variableSigma(weights.getNumVisible())
|
||||||
|
, m_progress(0)
|
||||||
|
{
|
||||||
|
Noise_Init(&m_noise, 0x32727155);
|
||||||
|
m_variableSigma.fill(m_params.m_constantSigma);
|
||||||
|
updateHiddenBatch();
|
||||||
|
}
|
||||||
|
|
||||||
|
Rbm::~Rbm()
|
||||||
|
{
|
||||||
|
Noise_Free(&m_noise);
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::sample(MatrixXd &srcDst)
|
||||||
|
{
|
||||||
|
sample(srcDst, srcDst);
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::sample(MatrixXd &dst, MatrixXd const &src)
|
||||||
|
{
|
||||||
|
uint32_t i;
|
||||||
|
|
||||||
|
for (i=0; i < src.array().size(); i++)
|
||||||
|
{
|
||||||
|
dst.array()(i) = (double)(src.array()(i) > Noise_Uniform(&m_noise));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::probsLogistic(MatrixXd &src)
|
||||||
|
{
|
||||||
|
src.array() = (-src.array()).exp();
|
||||||
|
src.array() += 1;
|
||||||
|
src.array() = 1.0/src.array();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::probsLogistic(RowVectorXd &src)
|
||||||
|
{
|
||||||
|
src.array() = (-src.array()).exp();
|
||||||
|
src.array() += 1;
|
||||||
|
src.array() = 1.0/src.array();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::probsLogistic(MatrixXd &src, const MatrixXd &sigma)
|
||||||
|
{
|
||||||
|
src.array() /= (sigma.array() + EPSILON_SIGMA);
|
||||||
|
src.array() = (-src.array()).exp();
|
||||||
|
src.array() += 1;
|
||||||
|
src.array() = 1.0/src.array();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::probsLogistic(RowVectorXd &src, const RowVectorXd &sigma)
|
||||||
|
{
|
||||||
|
src.array() /= (sigma.array() + EPSILON_SIGMA);
|
||||||
|
src.array() = (-src.array()).exp();
|
||||||
|
src.array() += 1;
|
||||||
|
src.array() = 1.0/src.array();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::probsGaussian(MatrixXd &src, const MatrixXd &sigma)
|
||||||
|
{
|
||||||
|
src.array() = 1 - src.array();
|
||||||
|
src.array() *= src.array();
|
||||||
|
src.array() *= -0.5;
|
||||||
|
|
||||||
|
MatrixXd var = sigma;
|
||||||
|
var.array() += EPSILON_SIGMA;
|
||||||
|
var.array() *= var.array();
|
||||||
|
|
||||||
|
src.array() /= var.array();
|
||||||
|
src.array() = src.array().exp();
|
||||||
|
|
||||||
|
MatrixXd k = var;
|
||||||
|
|
||||||
|
k.array() *= 2*3.14159265359;
|
||||||
|
k.array() = k.array().sqrt();
|
||||||
|
k.array() = 1.0/k.array();
|
||||||
|
|
||||||
|
src.array() *= k.array();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::probsGaussian(RowVectorXd &src, const RowVectorXd &sigma)
|
||||||
|
{
|
||||||
|
src.array() = 1 - src.array();
|
||||||
|
src.array() *= src.array();
|
||||||
|
src.array() *= -0.5;
|
||||||
|
|
||||||
|
RowVectorXd var = sigma;
|
||||||
|
var.array() += EPSILON_SIGMA;
|
||||||
|
var.array() *= var.array();
|
||||||
|
|
||||||
|
src.array() /= var.array();
|
||||||
|
src.array() = src.array().exp();
|
||||||
|
|
||||||
|
RowVectorXd k = var;
|
||||||
|
|
||||||
|
k.array() *= 2*3.14159265359;
|
||||||
|
k.array() = k.array().sqrt();
|
||||||
|
k.array() = 1.0/k.array();
|
||||||
|
|
||||||
|
src.array() *= k.array();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::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 Rbm::sampleGaussian(MatrixXd &src, const MatrixXd &sigma)
|
||||||
|
{
|
||||||
|
uint32_t i;
|
||||||
|
|
||||||
|
for (i=0; i < src.array().size(); i++)
|
||||||
|
{
|
||||||
|
src.array()(i) = sigma(i)*Noise_Gaussian(&m_noise) + src.array()(i);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RowVectorXd Rbm::normalizeData(RowVectorXd const &src, RowVectorXd const &mu, RowVectorXd const &var)
|
||||||
|
{
|
||||||
|
// Remove mean
|
||||||
|
RowVectorXd res = src - mu;
|
||||||
|
// res.array() /= var.array() + EPSILON_SIGMA;
|
||||||
|
|
||||||
|
// cout << __PRETTY_FUNCTION__ << ": " << res << endl;
|
||||||
|
return res;
|
||||||
|
}
|
||||||
|
|
||||||
|
RowVectorXd Rbm::calcMean(MatrixXd const &batch)
|
||||||
|
{
|
||||||
|
// Remove mean
|
||||||
|
RowVectorXd res = batch.colwise().mean();
|
||||||
|
|
||||||
|
// cout << __PRETTY_FUNCTION__ << ": " << res << endl;
|
||||||
|
return res;
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
RowVectorXd Rbm::calcSigma(MatrixXd const &batch)
|
||||||
|
{
|
||||||
|
MatrixXd x = batch.rowwise() - batch.colwise().mean();
|
||||||
|
|
||||||
|
x.array() *= x.array();
|
||||||
|
|
||||||
|
RowVectorXd res = x.colwise().mean().array().sqrt();
|
||||||
|
|
||||||
|
// cout << __PRETTY_FUNCTION__ << ": " << res << endl;
|
||||||
|
return res;
|
||||||
|
}
|
||||||
|
|
||||||
|
MatrixXd Rbm::calcZ(MatrixXd &v, MatrixXd &h)
|
||||||
|
{
|
||||||
|
|
||||||
|
MatrixXd t1(v.rows(), m_w.getNumVisible());
|
||||||
|
|
||||||
|
t1 = v - m_w.visibleBias().transpose().replicate(v.rows(), 1);
|
||||||
|
t1.array() *= t1.array();
|
||||||
|
t1.array() *= 0.5;
|
||||||
|
|
||||||
|
t1 -= (h * m_w.weights().transpose());
|
||||||
|
|
||||||
|
return t1;
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::train(uint32_t numEpochs, double sigmaMin)
|
||||||
|
{
|
||||||
|
uint32_t t, i;
|
||||||
|
uint32_t epoch;
|
||||||
|
uint32_t gibbs;
|
||||||
|
|
||||||
|
size_t batchSize = m_batch.rows();
|
||||||
|
|
||||||
|
double dProgress = 1.0/numEpochs;
|
||||||
|
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 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);
|
||||||
|
for (i=0; i < batchSize; i++)
|
||||||
|
{
|
||||||
|
RowVectorXd x = batch.row(i);
|
||||||
|
batch.row(i) = normalizeData(x, mean, m_variableSigma);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
m_progress = 0;
|
||||||
|
for (epoch=0; epoch < numEpochs; epoch++)
|
||||||
|
{
|
||||||
|
onProgressChanged();
|
||||||
|
|
||||||
|
if (m_params.m_doSampleBatch)
|
||||||
|
{
|
||||||
|
// When the hidden units are being driven by data, always use stochastic binary states
|
||||||
|
sample(batch_sampled, 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)
|
||||||
|
dBiasV_curr = batch.colwise().sum();
|
||||||
|
dBiasH_curr = h.colwise().sum();
|
||||||
|
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)
|
||||||
|
{
|
||||||
|
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_doSampleVisible)
|
||||||
|
{
|
||||||
|
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);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update weights (negative phase)
|
||||||
|
dBiasV_curr -= m_v.colwise().sum();
|
||||||
|
dBiasH_curr -= h.colwise().sum();
|
||||||
|
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;
|
||||||
|
|
||||||
|
if (m_params.m_doSparse)
|
||||||
|
{
|
||||||
|
MatrixXd h1 = h-MatrixXd::Ones(h.rows(), h.cols())*m_params.m_sparsity;
|
||||||
|
RowVectorXd hm = h1.colwise().mean();
|
||||||
|
m_w.hiddenBias() -= m_params.m_muSparsity * hm;
|
||||||
|
}
|
||||||
|
else
|
||||||
|
{
|
||||||
|
m_w.hiddenBias() += mu_biasH*(m_params.m_momentum*dBiasH + (1-m_params.m_momentum)*dBiasH_curr);
|
||||||
|
}
|
||||||
|
dBiasH = dBiasH_curr;
|
||||||
|
|
||||||
|
MatrixXd p = m_w.weights();
|
||||||
|
if (m_params.m_weightDecay > 0)
|
||||||
|
{
|
||||||
|
for (size_t row=0; row < m_w.weights().rows(); row++)
|
||||||
|
{
|
||||||
|
for (size_t col=0; col < m_w.weights().cols(); col++)
|
||||||
|
{
|
||||||
|
if (p(row, col) >= 0)
|
||||||
|
{
|
||||||
|
p(row, col) = m_params.m_weightDecay;
|
||||||
|
}
|
||||||
|
else
|
||||||
|
{
|
||||||
|
p(row, col) -= m_params.m_weightDecay;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
m_w.weights() -= mu_w*p;
|
||||||
|
}
|
||||||
|
m_w.weights() += mu_w*(m_params.m_momentum*dW + (1-m_params.m_momentum)*dW_curr);
|
||||||
|
dW = dW_curr;
|
||||||
|
|
||||||
|
if (m_variableSigma[0] > sigmaMin)
|
||||||
|
{
|
||||||
|
m_variableSigma.array() *= m_params.m_sigmaDecay;
|
||||||
|
}
|
||||||
|
|
||||||
|
m_progress += dProgress;
|
||||||
|
|
||||||
|
diffErr = m_batch - m_v;
|
||||||
|
diffErr.array() *= diffErr.array();
|
||||||
|
double err = diffErr.colwise().sum().sum();
|
||||||
|
|
||||||
|
cout << "err =" << endl;
|
||||||
|
cout << err << endl;
|
||||||
|
|
||||||
|
} // Number of epochs
|
||||||
|
|
||||||
|
updateHiddenBatch();
|
||||||
|
onProgressChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
double Rbm::getProgress() const
|
||||||
|
{
|
||||||
|
return m_progress;
|
||||||
|
}
|
||||||
|
|
||||||
|
double Rbm::getEnergy(const VectorXd& visible, const VectorXd& hidden)
|
||||||
|
{
|
||||||
|
double energy;
|
||||||
|
double sigma = m_variableSigma.array().mean();
|
||||||
|
|
||||||
|
energy = m_w.visibleBias() * visible;
|
||||||
|
energy += m_w.hiddenBias() * hidden;
|
||||||
|
energy += visible.transpose() * m_w.weights() * hidden;
|
||||||
|
|
||||||
|
return -energy/(sigma*sigma);
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::toHidden(RowVectorXd &h, RowVectorXd const &v)
|
||||||
|
{
|
||||||
|
h = v * m_w.weights();
|
||||||
|
h += m_w.hiddenBias();
|
||||||
|
probsLogistic(h);
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::toVisible(RowVectorXd &v, RowVectorXd const &h)
|
||||||
|
{
|
||||||
|
v = h * m_w.weights().transpose();
|
||||||
|
v += m_w.visibleBias();
|
||||||
|
|
||||||
|
if (m_params.m_useVisibleGaussian)
|
||||||
|
{
|
||||||
|
// probsGaussian(v, m_variableSigma);
|
||||||
|
}
|
||||||
|
else
|
||||||
|
{
|
||||||
|
probsLogistic(v, m_variableSigma);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setConstantSigma(double value)
|
||||||
|
{
|
||||||
|
m_params.m_constantSigma = value;
|
||||||
|
m_variableSigma.fill(m_params.m_constantSigma);
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
RowVectorXd& Rbm::getVariableSigma()
|
||||||
|
{
|
||||||
|
return m_variableSigma;
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setSigmaDecay(double value)
|
||||||
|
{
|
||||||
|
m_params.m_sigmaDecay = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setWeightDecay(double value)
|
||||||
|
{
|
||||||
|
m_params.m_weightDecay = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setLambda(double value)
|
||||||
|
{
|
||||||
|
m_params.m_lambda = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setSparsity(double value)
|
||||||
|
{
|
||||||
|
m_params.m_sparsity = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setUseVisibleGaussian(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_useVisibleGaussian = flag;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setDoRaoBlackwell(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_doRaoBlackwell = flag;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setDoSampleVisible(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_doSampleVisible = flag;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setDoSampleBatch(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_doSampleBatch = flag;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setDoSparse(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_doSparse = flag;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setNormalizeData(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_doNormalizeData = flag;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setDoLearnVariance(bool flag)
|
||||||
|
{
|
||||||
|
m_params.m_doLearnVariance = flag;
|
||||||
|
if (m_params.m_doLearnVariance && m_batch.rows())
|
||||||
|
{
|
||||||
|
m_variableSigma = calcSigma(m_batch);
|
||||||
|
}
|
||||||
|
else
|
||||||
|
{
|
||||||
|
m_variableSigma.fill(m_params.m_constantSigma);
|
||||||
|
}
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setNumGibbs(size_t value)
|
||||||
|
{
|
||||||
|
m_params.m_numGibbs = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setMuWeights(double value)
|
||||||
|
{
|
||||||
|
m_params.m_muWeights = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setMuSparsity(double value)
|
||||||
|
{
|
||||||
|
m_params.m_muSparsity = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::setMomentum(double value)
|
||||||
|
{
|
||||||
|
m_params.m_momentum = value;
|
||||||
|
onParamsChanged();
|
||||||
|
}
|
||||||
|
|
||||||
|
MatrixXd const& Rbm::getHiddenBatch()
|
||||||
|
{
|
||||||
|
return m_h;
|
||||||
|
}
|
||||||
|
|
||||||
|
MatrixXd const& Rbm::getVisibleBatch()
|
||||||
|
{
|
||||||
|
return m_v;
|
||||||
|
}
|
||||||
|
|
||||||
|
MatrixXd const& Rbm::getBatch()
|
||||||
|
{
|
||||||
|
return m_batch;
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::updateHiddenBatch()
|
||||||
|
{
|
||||||
|
m_h.resize(m_batch.rows(), m_w.getNumHidden());
|
||||||
|
toHiddenBatch(m_h, m_batch);
|
||||||
|
}
|
||||||
|
|
||||||
|
Rbm::Params const& Rbm::params()
|
||||||
|
{
|
||||||
|
return m_params;
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::toHiddenBatch(MatrixXd &h, MatrixXd const &v)
|
||||||
|
{
|
||||||
|
if (v.cols() == m_w.weights().rows())
|
||||||
|
{
|
||||||
|
h = v * m_w.weights();
|
||||||
|
h += m_w.hiddenBias().replicate(m_batch.rows(), 1);
|
||||||
|
probsLogistic(h);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
void Rbm::toVisibleBatch(MatrixXd &v, MatrixXd const &h)
|
||||||
|
{
|
||||||
|
v = h * m_w.weights().transpose();
|
||||||
|
v += m_w.visibleBias().replicate(m_batch.rows(), 1);
|
||||||
|
}
|
||||||
+46
-510
@@ -14,11 +14,6 @@
|
|||||||
|
|
||||||
using namespace Eigen;
|
using namespace Eigen;
|
||||||
|
|
||||||
void mylog(const char* format, ...);
|
|
||||||
#define printf mylog
|
|
||||||
|
|
||||||
#define EPSILON_SIGMA 0.001
|
|
||||||
|
|
||||||
class Rbm
|
class Rbm
|
||||||
{
|
{
|
||||||
public:
|
public:
|
||||||
@@ -27,7 +22,7 @@ public:
|
|||||||
Params()
|
Params()
|
||||||
: m_constantSigma(1.0)
|
: m_constantSigma(1.0)
|
||||||
, m_sigmaDecay(1.0)
|
, m_sigmaDecay(1.0)
|
||||||
, m_weightDecay(0.00001)
|
, m_weightDecay(0.0)
|
||||||
, m_lambda(1.0)
|
, m_lambda(1.0)
|
||||||
, m_sparsity(0.05)
|
, m_sparsity(0.05)
|
||||||
, m_muWeights(0.1)
|
, m_muWeights(0.1)
|
||||||
@@ -62,492 +57,49 @@ public:
|
|||||||
size_t m_numGibbs;
|
size_t m_numGibbs;
|
||||||
};
|
};
|
||||||
|
|
||||||
Rbm(Weights &weights, const MatrixXd &batch)
|
Rbm(Weights &weights, const MatrixXd &batch);
|
||||||
: m_w(weights)
|
~Rbm();
|
||||||
, m_batch(batch)
|
void sample(MatrixXd &srcDst);
|
||||||
, m_variableSigma(weights.getNumVisible())
|
void sample(MatrixXd &dst, MatrixXd const &src);
|
||||||
, m_progress(0)
|
static void probsLogistic(MatrixXd &src);
|
||||||
{
|
static void probsLogistic(RowVectorXd &src);
|
||||||
Noise_Init(&m_noise, 0x32727155);
|
static void probsLogistic(MatrixXd &src, const MatrixXd &sigma);
|
||||||
m_variableSigma.fill(m_params.m_constantSigma);
|
static void probsLogistic(RowVectorXd &src, const RowVectorXd &sigma);
|
||||||
updateHiddenBatch();
|
static void probsGaussian(MatrixXd &src, const MatrixXd &sigma);
|
||||||
}
|
static void probsGaussian(RowVectorXd &src, const RowVectorXd &sigma);
|
||||||
|
void sampleGaussian(MatrixXd &dst, MatrixXd const &src, const MatrixXd &sigma);
|
||||||
~Rbm()
|
void sampleGaussian(MatrixXd &src, const MatrixXd &sigma);
|
||||||
{
|
RowVectorXd normalizeData(RowVectorXd const &src, RowVectorXd const &mu, RowVectorXd const &var);
|
||||||
Noise_Free(&m_noise);
|
RowVectorXd calcMean(MatrixXd const &batch);
|
||||||
}
|
RowVectorXd calcSigma(MatrixXd const &batch);
|
||||||
|
MatrixXd calcZ(MatrixXd &v, MatrixXd &h);
|
||||||
void sample(MatrixXd &srcDst)
|
void train(uint32_t numEpochs, double sigmaMin = 0.05);
|
||||||
{
|
double getProgress() const;
|
||||||
sample(srcDst, srcDst);
|
double getEnergy(const VectorXd& visible, const VectorXd& hidden);
|
||||||
}
|
void toHidden(RowVectorXd &h, RowVectorXd const &v);
|
||||||
|
void toVisible(RowVectorXd &v, RowVectorXd const &h);
|
||||||
void sample(MatrixXd &dst, MatrixXd const &src)
|
void setConstantSigma(double value);
|
||||||
{
|
RowVectorXd& getVariableSigma();
|
||||||
uint32_t i;
|
void setSigmaDecay(double value);
|
||||||
|
void setWeightDecay(double value);
|
||||||
for (i=0; i < src.array().size(); i++)
|
void setLambda(double value);
|
||||||
{
|
void setSparsity(double value);
|
||||||
dst.array()(i) = (double)(src.array()(i) > Noise_Uniform(&m_noise));
|
void setUseVisibleGaussian(bool flag);
|
||||||
}
|
void setDoRaoBlackwell(bool flag);
|
||||||
}
|
void setDoSampleVisible(bool flag);
|
||||||
|
void setDoSampleBatch(bool flag);
|
||||||
void probsLogistic(MatrixXd &src)
|
void setDoSparse(bool flag);
|
||||||
{
|
void setNormalizeData(bool flag);
|
||||||
src.array() = (-src.array()).exp();
|
void setDoLearnVariance(bool flag);
|
||||||
src.array() += 1;
|
void setNumGibbs(size_t value);
|
||||||
src.array() = 1.0/src.array();
|
void setMuWeights(double value);
|
||||||
}
|
void setMuSparsity(double value);
|
||||||
|
void setMomentum(double value);
|
||||||
void probsLogistic(RowVectorXd &src)
|
MatrixXd const& getHiddenBatch();
|
||||||
{
|
MatrixXd const& getVisibleBatch();
|
||||||
src.array() = (-src.array()).exp();
|
MatrixXd const& getBatch();
|
||||||
src.array() += 1;
|
void updateHiddenBatch();
|
||||||
src.array() = 1.0/src.array();
|
Params const& params();
|
||||||
}
|
|
||||||
|
|
||||||
void probsLogistic(MatrixXd &src, const MatrixXd &sigma)
|
|
||||||
{
|
|
||||||
src.array() /= (sigma.array() + EPSILON_SIGMA);
|
|
||||||
src.array() = (-src.array()).exp();
|
|
||||||
src.array() += 1;
|
|
||||||
src.array() = 1.0/src.array();
|
|
||||||
}
|
|
||||||
|
|
||||||
void probsLogistic(RowVectorXd &src, const RowVectorXd &sigma)
|
|
||||||
{
|
|
||||||
src.array() /= (sigma.array() + EPSILON_SIGMA);
|
|
||||||
src.array() = (-src.array()).exp();
|
|
||||||
src.array() += 1;
|
|
||||||
src.array() = 1.0/src.array();
|
|
||||||
}
|
|
||||||
|
|
||||||
void probsGaussian(MatrixXd &src, const MatrixXd &sigma)
|
|
||||||
{
|
|
||||||
src.array() = 1 - src.array();
|
|
||||||
src.array() *= src.array();
|
|
||||||
src.array() *= -0.5;
|
|
||||||
|
|
||||||
MatrixXd var = sigma;
|
|
||||||
var.array() += EPSILON_SIGMA;
|
|
||||||
var.array() *= var.array();
|
|
||||||
|
|
||||||
src.array() /= var.array();
|
|
||||||
src.array() = src.array().exp();
|
|
||||||
|
|
||||||
MatrixXd k = var;
|
|
||||||
|
|
||||||
k.array() *= 2*3.14159265359;
|
|
||||||
k.array() = k.array().sqrt();
|
|
||||||
k.array() = 1.0/k.array();
|
|
||||||
|
|
||||||
src.array() *= k.array();
|
|
||||||
}
|
|
||||||
|
|
||||||
void probsGaussian(RowVectorXd &src, const RowVectorXd &sigma)
|
|
||||||
{
|
|
||||||
src.array() = 1 - src.array();
|
|
||||||
src.array() *= src.array();
|
|
||||||
src.array() *= -0.5;
|
|
||||||
|
|
||||||
RowVectorXd var = sigma;
|
|
||||||
var.array() += EPSILON_SIGMA;
|
|
||||||
var.array() *= var.array();
|
|
||||||
|
|
||||||
src.array() /= var.array();
|
|
||||||
src.array() = src.array().exp();
|
|
||||||
|
|
||||||
RowVectorXd k = var;
|
|
||||||
|
|
||||||
k.array() *= 2*3.14159265359;
|
|
||||||
k.array() = k.array().sqrt();
|
|
||||||
k.array() = 1.0/k.array();
|
|
||||||
|
|
||||||
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;
|
|
||||||
|
|
||||||
for (i=0; i < src.array().size(); i++)
|
|
||||||
{
|
|
||||||
src.array()(i) = sigma(i)*Noise_Gaussian(&m_noise) + src.array()(i);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RowVectorXd normalizeData(RowVectorXd const &src, RowVectorXd const &mu, RowVectorXd const &var)
|
|
||||||
{
|
|
||||||
// Remove mean
|
|
||||||
RowVectorXd res = src - mu;
|
|
||||||
// res.array() /= var.array() + EPSILON_SIGMA;
|
|
||||||
|
|
||||||
// cout << __PRETTY_FUNCTION__ << ": " << res << endl;
|
|
||||||
return res;
|
|
||||||
}
|
|
||||||
|
|
||||||
RowVectorXd calcMean(MatrixXd const &batch)
|
|
||||||
{
|
|
||||||
// Remove mean
|
|
||||||
RowVectorXd res = batch.colwise().mean();
|
|
||||||
|
|
||||||
// cout << __PRETTY_FUNCTION__ << ": " << res << endl;
|
|
||||||
return res;
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
RowVectorXd calcSigma(MatrixXd const &batch)
|
|
||||||
{
|
|
||||||
MatrixXd x = batch.rowwise() - batch.colwise().mean();
|
|
||||||
|
|
||||||
x.array() *= x.array();
|
|
||||||
|
|
||||||
RowVectorXd res = x.colwise().mean().array().sqrt();
|
|
||||||
|
|
||||||
// cout << __PRETTY_FUNCTION__ << ": " << res << endl;
|
|
||||||
return res;
|
|
||||||
}
|
|
||||||
|
|
||||||
MatrixXd calcZ(MatrixXd &v, MatrixXd &h)
|
|
||||||
{
|
|
||||||
|
|
||||||
MatrixXd t1(v.rows(), m_w.getNumVisible());
|
|
||||||
|
|
||||||
t1 = v - m_w.visibleBias().transpose().replicate(v.rows(), 1);
|
|
||||||
t1.array() *= t1.array();
|
|
||||||
t1.array() *= 0.5;
|
|
||||||
|
|
||||||
t1 -= (h * m_w.weights().transpose());
|
|
||||||
|
|
||||||
return t1;
|
|
||||||
}
|
|
||||||
|
|
||||||
void train(uint32_t numEpochs, double sigmaMin = 0.05)
|
|
||||||
{
|
|
||||||
uint32_t t, i;
|
|
||||||
uint32_t epoch;
|
|
||||||
uint32_t gibbs;
|
|
||||||
|
|
||||||
size_t batchSize = m_batch.rows();
|
|
||||||
|
|
||||||
double dProgress = 1.0/numEpochs;
|
|
||||||
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 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);
|
|
||||||
for (i=0; i < batchSize; i++)
|
|
||||||
{
|
|
||||||
RowVectorXd x = batch.row(i);
|
|
||||||
batch.row(i) = normalizeData(x, mean, m_variableSigma);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
m_progress = 0;
|
|
||||||
for (epoch=0; epoch < numEpochs; epoch++)
|
|
||||||
{
|
|
||||||
onProgressChanged();
|
|
||||||
|
|
||||||
if (m_params.m_doSampleBatch)
|
|
||||||
{
|
|
||||||
// When the hidden units are being driven by data, always use stochastic binary states
|
|
||||||
sample(batch_sampled, 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)
|
|
||||||
dBiasV_curr = batch.colwise().sum();
|
|
||||||
dBiasH_curr = h.colwise().sum();
|
|
||||||
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)
|
|
||||||
{
|
|
||||||
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_doSampleVisible)
|
|
||||||
{
|
|
||||||
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);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Update weights (negative phase)
|
|
||||||
dBiasV_curr -= m_v.colwise().sum();
|
|
||||||
dBiasH_curr -= h.colwise().sum();
|
|
||||||
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;
|
|
||||||
|
|
||||||
if (m_params.m_doSparse)
|
|
||||||
{
|
|
||||||
MatrixXd h1 = h-MatrixXd::Ones(h.rows(), h.cols())*m_params.m_sparsity;
|
|
||||||
RowVectorXd hm = h1.colwise().mean();
|
|
||||||
m_w.hiddenBias() -= m_params.m_muSparsity * hm;
|
|
||||||
}
|
|
||||||
else
|
|
||||||
{
|
|
||||||
m_w.hiddenBias() += mu_biasH*(m_params.m_momentum*dBiasH + (1-m_params.m_momentum)*dBiasH_curr);
|
|
||||||
}
|
|
||||||
dBiasH = dBiasH_curr;
|
|
||||||
|
|
||||||
m_w.weights() -= m_params.m_weightDecay*m_w.weights();
|
|
||||||
m_w.weights() += mu_w*(m_params.m_momentum*dW + (1-m_params.m_momentum)*dW_curr);
|
|
||||||
dW = dW_curr;
|
|
||||||
|
|
||||||
if (m_variableSigma[0] > sigmaMin)
|
|
||||||
{
|
|
||||||
m_variableSigma.array() *= m_params.m_sigmaDecay;
|
|
||||||
}
|
|
||||||
|
|
||||||
m_progress += dProgress;
|
|
||||||
|
|
||||||
diffErr = m_batch - m_v;
|
|
||||||
diffErr.array() *= diffErr.array();
|
|
||||||
double err = diffErr.colwise().sum().sum();
|
|
||||||
|
|
||||||
cout << "err =" << endl;
|
|
||||||
cout << err << endl;
|
|
||||||
|
|
||||||
} // Number of epochs
|
|
||||||
|
|
||||||
updateHiddenBatch();
|
|
||||||
onProgressChanged();
|
|
||||||
}
|
|
||||||
|
|
||||||
double getProgress() const
|
|
||||||
{
|
|
||||||
return m_progress;
|
|
||||||
}
|
|
||||||
|
|
||||||
double getEnergy(const VectorXd& visible, const VectorXd& hidden)
|
|
||||||
{
|
|
||||||
double energy;
|
|
||||||
double sigma = m_variableSigma.array().mean();
|
|
||||||
|
|
||||||
energy = m_w.visibleBias() * visible;
|
|
||||||
energy += m_w.hiddenBias() * hidden;
|
|
||||||
energy += visible.transpose() * m_w.weights() * hidden;
|
|
||||||
|
|
||||||
return -energy/(sigma*sigma);
|
|
||||||
}
|
|
||||||
|
|
||||||
void toHidden(RowVectorXd &h, RowVectorXd const &v)
|
|
||||||
{
|
|
||||||
h = v * m_w.weights();
|
|
||||||
h += m_w.hiddenBias();
|
|
||||||
probsLogistic(h);
|
|
||||||
}
|
|
||||||
|
|
||||||
void toVisible(RowVectorXd &v, RowVectorXd const &h)
|
|
||||||
{
|
|
||||||
v = h * m_w.weights().transpose();
|
|
||||||
v += m_w.visibleBias();
|
|
||||||
|
|
||||||
if (m_params.m_useVisibleGaussian)
|
|
||||||
{
|
|
||||||
probsGaussian(v, m_variableSigma);
|
|
||||||
}
|
|
||||||
else
|
|
||||||
{
|
|
||||||
probsLogistic(v, m_variableSigma);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
void setConstantSigma(double value)
|
|
||||||
{
|
|
||||||
m_params.m_constantSigma = value;
|
|
||||||
m_variableSigma.fill(m_params.m_constantSigma);
|
|
||||||
onParamsChanged();
|
|
||||||
}
|
|
||||||
|
|
||||||
RowVectorXd& getVariableSigma()
|
|
||||||
{
|
|
||||||
return m_variableSigma;
|
|
||||||
}
|
|
||||||
|
|
||||||
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 setDoSampleVisible(bool 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)
|
|
||||||
{
|
|
||||||
m_params.m_doLearnVariance = flag;
|
|
||||||
if (m_params.m_doLearnVariance && m_batch.rows())
|
|
||||||
{
|
|
||||||
m_variableSigma = calcSigma(m_batch);
|
|
||||||
}
|
|
||||||
else
|
|
||||||
{
|
|
||||||
m_variableSigma.fill(m_params.m_constantSigma);
|
|
||||||
}
|
|
||||||
onParamsChanged();
|
|
||||||
}
|
|
||||||
|
|
||||||
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()
|
|
||||||
{
|
|
||||||
return m_h;
|
|
||||||
}
|
|
||||||
|
|
||||||
MatrixXd const& getVisibleBatch()
|
|
||||||
{
|
|
||||||
return m_v;
|
|
||||||
}
|
|
||||||
|
|
||||||
MatrixXd const& getBatch()
|
|
||||||
{
|
|
||||||
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:
|
private:
|
||||||
Weights &m_w;
|
Weights &m_w;
|
||||||
@@ -559,21 +111,8 @@ private:
|
|||||||
double m_progress;
|
double m_progress;
|
||||||
Params m_params;
|
Params m_params;
|
||||||
|
|
||||||
void toHiddenBatch(MatrixXd &h, MatrixXd const &v)
|
void toHiddenBatch(MatrixXd &h, MatrixXd const &v);
|
||||||
{
|
void toVisibleBatch(MatrixXd &v, MatrixXd const &h);
|
||||||
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)
|
|
||||||
{
|
|
||||||
v = h * m_w.weights().transpose();
|
|
||||||
v += m_w.visibleBias().replicate(m_batch.rows(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
protected:
|
protected:
|
||||||
virtual void onProgressChanged()
|
virtual void onProgressChanged()
|
||||||
@@ -585,7 +124,4 @@ protected:
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#endif /* RBM_HPP_ */
|
#endif /* RBM_HPP_ */
|
||||||
|
|||||||
Reference in New Issue
Block a user