- fully refactored traning
- removed non working stuff git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@553 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+84
-210
@@ -60,47 +60,42 @@ void Rbm::sampleGaussian(MatrixXd &srcDst)
|
||||
sampleGaussian(srcDst, srcDst);
|
||||
}
|
||||
|
||||
void Rbm::sample(MatrixXd &srcDst)
|
||||
{
|
||||
sample(srcDst, srcDst);
|
||||
}
|
||||
|
||||
void Rbm::sample(MatrixXd &dst, MatrixXd const &src)
|
||||
MatrixXd Rbm::sample(MatrixXd const &src)
|
||||
{
|
||||
MatrixXd n(src.rows(), src.cols());
|
||||
|
||||
noiseUniform(n);
|
||||
|
||||
dst = (src.array() > n.array()).cast<double>();
|
||||
return (src.array() > n.array()).cast<double>();
|
||||
}
|
||||
|
||||
void Rbm::sample(MatrixXd &dst, MatrixXd const &src)
|
||||
{
|
||||
dst = sample(src);
|
||||
}
|
||||
|
||||
MatrixXd Rbm::probsLogistic(MatrixXd const &src)
|
||||
{
|
||||
return (1 + (-src.array()).exp()).array().cwiseInverse();
|
||||
}
|
||||
|
||||
RowVectorXd Rbm::probsLogistic(RowVectorXd const &src)
|
||||
{
|
||||
return (1 + (-src.array()).exp()).array().cwiseInverse();
|
||||
}
|
||||
|
||||
MatrixXd Rbm::normalizeData(MatrixXd const &src)
|
||||
{
|
||||
double mean = src.array().mean();
|
||||
cout << "mean" << ": " << endl << mean << endl;
|
||||
|
||||
MatrixXd x = src - mean*MatrixXd::Ones(src.rows(), src.cols());
|
||||
MatrixXd x2 = x.cwiseProduct(x);
|
||||
double stddev = sqrt(x2.array().mean());
|
||||
cout << "stddev" << ": " << endl << stddev << endl;
|
||||
|
||||
return x;
|
||||
|
||||
}
|
||||
|
||||
void Rbm::probsLogistic(MatrixXd &srcDst)
|
||||
{
|
||||
srcDst = (1 + (-srcDst.array()).exp()).array().cwiseInverse();
|
||||
}
|
||||
|
||||
void Rbm::probsLogistic(RowVectorXd &srcDst)
|
||||
{
|
||||
srcDst = (1 + (-srcDst.array()).exp()).array().cwiseInverse();
|
||||
}
|
||||
|
||||
void Rbm::normalizeData(MatrixXd &dst, MatrixXd const &src)
|
||||
{
|
||||
MatrixXd mean = src.rowwise().mean();
|
||||
// cout << "mean" << ": " << endl << mean << endl;
|
||||
|
||||
dst = src - mean.replicate(1, src.cols());
|
||||
// cout << "dst - mean" << ": " << endl << dst << endl;
|
||||
|
||||
MatrixXd x = dst.array().square();
|
||||
|
||||
MatrixXd var = x.rowwise().mean();
|
||||
// cout << "var" << ": " << endl << var << endl;
|
||||
|
||||
MatrixXd stddev_norm = var.array().sqrt().cwiseInverse();
|
||||
dst.array() *= stddev_norm.replicate(1, src.cols()).array();
|
||||
// cout << "dst" << ": " << endl << dst << endl;
|
||||
}
|
||||
|
||||
void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
@@ -118,7 +113,7 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
MatrixXd grad_bias_v(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||
MatrixXd grad_bias_h(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||
MatrixXd grad_weight(MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden()));
|
||||
MatrixXd __batch = m_batch_normalized;
|
||||
MatrixXd __batch = m_batch;
|
||||
MatrixXd momentum_weights = MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden());
|
||||
MatrixXd momentum_bias_v(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||
MatrixXd momentum_bias_h(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||
@@ -126,11 +121,6 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
double L1 = 0;
|
||||
double L2 = 0;
|
||||
|
||||
if (m_params.m_doNormalizeData && !m_params.m_useVisibleGaussian)
|
||||
{
|
||||
probsLogistic(__batch);
|
||||
}
|
||||
|
||||
m_progress = 0;
|
||||
while (trainingSizeRemain)
|
||||
{
|
||||
@@ -140,94 +130,75 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
trainingSizeRemain -= toSlice;
|
||||
batchRowIndex += toSlice;
|
||||
size_t batchSize = batch.rows();
|
||||
double learning_rate = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||
double learning_rate = m_params.m_learningRate/std::min(miniBatchSize, trainingSize);
|
||||
double weight_decay = m_params.m_weightDecay/std::min(miniBatchSize, trainingSize);
|
||||
|
||||
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
||||
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
||||
MatrixXd vis(batchSize, m_w.getNumVisible());
|
||||
MatrixXd hid(batchSize, m_w.getNumHidden());
|
||||
MatrixXd vis_state(batchSize, m_w.getNumVisible());
|
||||
MatrixXd vis_probs(batchSize, m_w.getNumVisible());
|
||||
MatrixXd hid_state(batchSize, m_w.getNumHidden());
|
||||
MatrixXd hid_probs(batchSize, m_w.getNumHidden());
|
||||
|
||||
for (epoch=0; epoch < numEpochs; epoch++)
|
||||
{
|
||||
onProgressChanged();
|
||||
|
||||
// Create hidden layer base on training data
|
||||
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
|
||||
hid = batch_sampled * m_w.weights();
|
||||
hid += m_w.hiddenBias().replicate(batchSize, 1);
|
||||
probsLogistic(hid);
|
||||
vis_state = sample(batch);
|
||||
}
|
||||
else
|
||||
{
|
||||
// Create hidden layer base on training data
|
||||
hid = batch * m_w.weights();
|
||||
hid += m_w.hiddenBias().replicate(batchSize, 1);
|
||||
probsLogistic(hid);
|
||||
vis_state = batch;
|
||||
}
|
||||
|
||||
hid_state = vis_state * m_w.weights() + m_w.hiddenBias().replicate(batchSize, 1);
|
||||
hid_probs = probsLogistic(hid_state);
|
||||
|
||||
// Sample hidden
|
||||
if (!m_params.m_doRaoBlackwell)
|
||||
if (m_params.m_doRaoBlackwell)
|
||||
{
|
||||
sample(hid);
|
||||
hid_state = hid_probs;
|
||||
}
|
||||
else
|
||||
{
|
||||
hid_state = sample(hid_probs);
|
||||
}
|
||||
|
||||
// Update weights (positive phase)
|
||||
grad_bias_v = batch.colwise().sum();
|
||||
grad_bias_h = hid.colwise().sum();
|
||||
grad_weight = batch.transpose() * hid;
|
||||
grad_weight = vis_state.transpose() * hid_state;
|
||||
grad_bias_v = vis_state.colwise().sum();
|
||||
grad_bias_h = hid_state.colwise().sum();
|
||||
|
||||
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
|
||||
{
|
||||
if (m_params.m_useHiddenGaussian)
|
||||
|
||||
// Create hidden representation given v
|
||||
hid_state = sample(hid_probs);
|
||||
|
||||
// Create visible reconstruction (a fantasy...) given hid
|
||||
vis_state = hid_state * m_w.weights().transpose() + m_w.visibleBias().replicate(batchSize, 1);
|
||||
vis_probs = probsLogistic(vis_state);
|
||||
if (m_params.m_doSampleVisible)
|
||||
{
|
||||
sampleGaussian(hid);
|
||||
vis_state = sample(vis_probs);
|
||||
}
|
||||
else
|
||||
{
|
||||
sample(hid);
|
||||
}
|
||||
if (m_params.m_useVisibleGaussian)
|
||||
{
|
||||
// Create visible reconstruction (a fantasy...) given hid
|
||||
vis = hid * m_w.weights().transpose();
|
||||
vis += m_w.visibleBias().replicate(batchSize, 1);
|
||||
sampleGaussian(v_sampled, vis);
|
||||
hid = v_sampled * m_w.weights();
|
||||
hid += m_w.hiddenBias().replicate(batchSize, 1);
|
||||
probsLogistic(hid);
|
||||
|
||||
}
|
||||
else
|
||||
{
|
||||
// Create visible reconstruction (a fantasy...) given hid
|
||||
vis = hid * m_w.weights().transpose();
|
||||
vis += m_w.visibleBias().replicate(batchSize, 1);
|
||||
probsLogistic(vis);
|
||||
if (m_params.m_doSampleVisible)
|
||||
{
|
||||
sample(v_sampled, vis);
|
||||
// Create hidden representation given sampled v
|
||||
hid = v_sampled * m_w.weights();
|
||||
hid += m_w.hiddenBias().replicate(batchSize, 1);
|
||||
probsLogistic(hid);
|
||||
}
|
||||
else
|
||||
{
|
||||
// Create hidden representation given v
|
||||
hid = vis * m_w.weights();
|
||||
hid += m_w.hiddenBias().replicate(batchSize, 1);
|
||||
probsLogistic(hid);
|
||||
}
|
||||
vis_state = vis_probs;
|
||||
}
|
||||
|
||||
// Create hidden representation given v
|
||||
hid_state = vis_state * m_w.weights() + m_w.hiddenBias().replicate(batchSize, 1);
|
||||
hid_probs = probsLogistic(hid_state);
|
||||
|
||||
}
|
||||
|
||||
// Update weights (negative phase)
|
||||
grad_bias_v -= vis.colwise().sum();
|
||||
grad_bias_h -= hid.colwise().sum();
|
||||
grad_weight -= vis.transpose() * hid;
|
||||
grad_bias_v -= vis_probs.colwise().sum();
|
||||
grad_bias_h -= hid_probs.colwise().sum();
|
||||
grad_weight -= vis_probs.transpose() * hid_probs;
|
||||
|
||||
for (int i=0; i < m_w.weights().rows(); i++)
|
||||
{
|
||||
@@ -235,11 +206,11 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
{
|
||||
if (m_w.weights()(i,j) >= 0)
|
||||
{
|
||||
penalty_weights(i,j) = m_params.m_weightDecay;
|
||||
penalty_weights(i,j) = weight_decay;
|
||||
}
|
||||
else
|
||||
{
|
||||
penalty_weights(i,j) = -m_params.m_weightDecay;
|
||||
penalty_weights(i,j) = -weight_decay;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -251,34 +222,16 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
momentum_weights = m_params.m_momentum*momentum_weights + grad_weight - L2*penalty_weights;
|
||||
|
||||
m_w.visibleBias() += learning_rate*momentum_bias_v;
|
||||
|
||||
if (m_params.m_doSparse)
|
||||
{
|
||||
MatrixXd h1 = hid-MatrixXd::Ones(hid.rows(), hid.cols())*m_params.m_sparsity;
|
||||
RowVectorXd hm = h1.colwise().mean();
|
||||
m_w.hiddenBias() -= m_params.m_muSparsity * hm;
|
||||
}
|
||||
else
|
||||
{
|
||||
m_w.hiddenBias() += learning_rate*momentum_bias_h;
|
||||
}
|
||||
|
||||
m_w.hiddenBias() += learning_rate*momentum_bias_h;
|
||||
m_w.weights() += learning_rate*momentum_weights;
|
||||
|
||||
if (m_params.m_sigmaDecay > 0)
|
||||
{
|
||||
if (m_variableSigma[0] > sigmaMin)
|
||||
{
|
||||
}
|
||||
}
|
||||
|
||||
m_progress += dProgress;
|
||||
|
||||
} // Number of epochs
|
||||
|
||||
MatrixXd diffErr = batch - vis;
|
||||
MatrixXd diffErr = batch - vis_probs;
|
||||
diffErr.array() *= diffErr.array();
|
||||
double err = diffErr.colwise().sum().sum();
|
||||
double err = diffErr.colwise().sum().mean();
|
||||
cout << "error (per mini batch) = " << err << endl;
|
||||
cout << "L1 = " << L1 << endl;
|
||||
cout << "L2 = " << L2 << endl;
|
||||
@@ -288,10 +241,9 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||
|
||||
MatrixXd vis = m_h * m_w.weights().transpose();
|
||||
vis += m_w.visibleBias().replicate(__batch.rows(), 1);
|
||||
probsLogistic(vis);
|
||||
MatrixXd diffErr = __batch - vis;
|
||||
MatrixXd diffErr = __batch - probsLogistic(vis);
|
||||
diffErr.array() *= diffErr.array();
|
||||
double err = diffErr.colwise().sum().sum();
|
||||
double err = diffErr.colwise().sum().mean();
|
||||
cout << "error (total) = " << err << endl;
|
||||
|
||||
onProgressChanged();
|
||||
@@ -307,35 +259,14 @@ void Rbm::toHidden(RowVectorXd &h, RowVectorXd const &v)
|
||||
{
|
||||
h = v * m_w.weights();
|
||||
h += m_w.hiddenBias();
|
||||
probsLogistic(h);
|
||||
h = 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)
|
||||
{
|
||||
probsLogistic(v);
|
||||
}
|
||||
}
|
||||
|
||||
void Rbm::setConstantSigma(double value)
|
||||
{
|
||||
m_params.m_constantSigma = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
RowVectorXd& Rbm::getVariableSigma()
|
||||
{
|
||||
return m_variableSigma;
|
||||
}
|
||||
|
||||
void Rbm::setSigmaDecay(double value)
|
||||
{
|
||||
m_params.m_sigmaDecay = value;
|
||||
onParamsChanged();
|
||||
v = probsLogistic(v);
|
||||
}
|
||||
|
||||
void Rbm::setWeightDecay(double value)
|
||||
@@ -344,30 +275,6 @@ void Rbm::setWeightDecay(double 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::setUseHiddenGaussian(bool flag)
|
||||
{
|
||||
m_params.m_useHiddenGaussian = flag;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void Rbm::setDoRaoBlackwell(bool flag)
|
||||
{
|
||||
m_params.m_doRaoBlackwell = flag;
|
||||
@@ -386,24 +293,6 @@ void Rbm::setDoSampleBatch(bool 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;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void Rbm::setNumGibbs(size_t value)
|
||||
{
|
||||
m_params.m_numGibbs = value;
|
||||
@@ -412,13 +301,7 @@ void Rbm::setNumGibbs(size_t value)
|
||||
|
||||
void Rbm::setMuWeights(double value)
|
||||
{
|
||||
m_params.m_muWeights = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
void Rbm::setMuSparsity(double value)
|
||||
{
|
||||
m_params.m_muSparsity = value;
|
||||
m_params.m_learningRate = value;
|
||||
onParamsChanged();
|
||||
}
|
||||
|
||||
@@ -435,7 +318,7 @@ MatrixXd const& Rbm::getHiddenBatch()
|
||||
|
||||
MatrixXd const& Rbm::getBatch()
|
||||
{
|
||||
return m_batch_normalized;
|
||||
return m_batch;
|
||||
}
|
||||
|
||||
void Rbm::updateHiddenBatch()
|
||||
@@ -444,20 +327,11 @@ void Rbm::updateHiddenBatch()
|
||||
{
|
||||
return;
|
||||
}
|
||||
m_batch_normalized.resize(m_batch.rows(), m_w.getNumHidden());
|
||||
|
||||
if (m_params.m_doNormalizeData)
|
||||
{
|
||||
normalizeData(m_batch_normalized, m_batch);
|
||||
}
|
||||
else
|
||||
{
|
||||
m_batch_normalized = m_batch;
|
||||
}
|
||||
m_h.resize(m_batch.rows(), m_w.getNumHidden());
|
||||
m_h = m_batch_normalized * m_w.weights();
|
||||
m_h = m_batch * m_w.weights();
|
||||
m_h += m_w.hiddenBias().replicate(m_batch.rows(), 1);
|
||||
probsLogistic(m_h);
|
||||
m_h = probsLogistic(m_h);
|
||||
}
|
||||
|
||||
Rbm::Params const& Rbm::params()
|
||||
|
||||
Reference in New Issue
Block a user