[RBM]
- 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:
+142
-133
@@ -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
@@ -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()
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user