- introduced mini batch training


git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@304 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-06-30 07:52:49 +00:00
parent 6e93c07863
commit 4979ab41b4
2 changed files with 143 additions and 137 deletions
+142 -133
View File
@@ -184,20 +184,16 @@ MatrixXd Rbm::calcZ(MatrixXd &v, MatrixXd &h)
void Rbm::train(uint32_t numEpochs, double sigmaMin)
{
uint32_t t, i;
uint32_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());
size_t trainingSize = m_batch.rows();
size_t trainingSizeRemain = trainingSize;
size_t batchRowIndex = 0;
const size_t miniBatchSize = 100;
double dProgress = 1.0/(numEpochs*(double)trainingSize/miniBatchSize);
MatrixXd dBiasV_curr(MatrixXd::Zero(1, m_w.getNumVisible()));
MatrixXd dBiasH_curr(MatrixXd::Zero(1, m_w.getNumHidden()));
@@ -206,137 +202,164 @@ void Rbm::train(uint32_t numEpochs, double sigmaMin)
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++)
while (trainingSizeRemain)
{
onProgressChanged();
cout << "trainingSizeRemain: " << trainingSizeRemain << endl;
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
MatrixXd batch = m_batch.block(batchRowIndex, 0, toSlice, m_w.getNumVisible());
trainingSizeRemain -= toSlice;
size_t batchSize = batch.rows();
double mu_w = m_params.m_muWeights/batchSize;
double mu_biasV = m_params.m_muWeights/batchSize;
double mu_biasH = m_params.m_muWeights/batchSize;
MatrixXd diffErr(batchSize, m_w.getNumVisible());
if (m_params.m_doSampleBatch)
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());
if (m_params.m_doNormalizeData)
{
// 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)
RowVectorXd mean = calcMean(batch);
for (i=0; i < batchSize; i++)
{
sampleGaussian(v_sampled, m_v, m_variableSigma.replicate(batchSize, 1));
toHiddenBatch(h, v_sampled);
RowVectorXd x = batch.row(i);
batch.row(i) = normalizeData(x, mean, m_variableSigma);
}
}
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
hid = batch_sampled * m_w.weights();
hid += m_w.hiddenBias().replicate(batchSize, 1);
probsLogistic(hid);
}
else
{
probsLogistic(m_v, m_variableSigma.replicate(batchSize, 1));
if (m_params.m_doSampleVisible)
// Create hidden layer base on training data
hid = batch * m_w.weights();
hid += m_w.hiddenBias().replicate(batchSize, 1);
probsLogistic(hid);
}
// Sample hidden
if (!m_params.m_doRaoBlackwell)
{
sample(hid);
}
// Update weights (positive phase)
dBiasV_curr = batch.colwise().sum();
dBiasH_curr = hid.colwise().sum();
dW_curr = batch.transpose() * hid;
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
{
sample(hid);
// Create visible reconstruction (a fantasy...) given hid
vis = hid * m_w.weights().transpose();
vis += m_w.visibleBias().replicate(batchSize, 1);
if (m_params.m_useVisibleGaussian)
{
sample(v_sampled, m_v);
// Create hidden representation given sampled v
toHiddenBatch(h, v_sampled);
sampleGaussian(v_sampled, vis, m_variableSigma.replicate(batchSize, 1));
hid = v_sampled * m_w.weights();
hid += m_w.hiddenBias().replicate(batchSize, 1);
probsLogistic(hid);
}
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)
probsLogistic(vis, m_variableSigma.replicate(batchSize, 1));
if (m_params.m_doSampleVisible)
{
p(row, col) = m_params.m_weightDecay;
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
{
p(row, col) -= m_params.m_weightDecay;
// Create hidden representation given v
hid = vis * m_w.weights();
hid += m_w.hiddenBias().replicate(batchSize, 1);
probsLogistic(hid);
}
}
}
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;
}
// Update weights (negative phase)
dBiasV_curr -= vis.colwise().sum();
dBiasH_curr -= hid.colwise().sum();
dW_curr -= vis.transpose() * hid;
m_progress += dProgress;
m_w.visibleBias() += mu_biasV*(m_params.m_momentum*dBiasV + (1-m_params.m_momentum)*dBiasV_curr);
dBiasV = dBiasV_curr;
diffErr = m_batch - m_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() += 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_params.m_sigmaDecay > 0)
{
if (m_variableSigma[0] > sigmaMin)
{
m_variableSigma.array() *= (1-m_params.m_sigmaDecay);
}
}
m_progress += dProgress;
} // Number of epochs
diffErr = batch - vis;
diffErr.array() *= diffErr.array();
double err = diffErr.colwise().sum().sum();
cout << "err =" << endl;
cout << err << endl;
} // Number of epochs
} // number of mini batches
updateHiddenBatch();
onProgressChanged();
}
@@ -508,26 +531,12 @@ MatrixXd const& Rbm::getBatch()
void Rbm::updateHiddenBatch()
{
m_h.resize(m_batch.rows(), m_w.getNumHidden());
toHiddenBatch(m_h, m_batch);
m_h = m_batch * m_w.weights();
m_h += m_w.hiddenBias().replicate(m_batch.rows(), 1);
probsLogistic(m_h);
}
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);
}
+1 -4
View File
@@ -21,7 +21,7 @@ public:
{
Params()
: m_constantSigma(1.0)
, m_sigmaDecay(1.0)
, m_sigmaDecay(0.0)
, m_weightDecay(0.0)
, m_lambda(1.0)
, m_sparsity(0.05)
@@ -111,9 +111,6 @@ private:
double m_progress;
Params m_params;
void toHiddenBatch(MatrixXd &h, MatrixXd const &v);
void toVisibleBatch(MatrixXd &v, MatrixXd const &h);
protected:
virtual void onProgressChanged()
{