- refactored

- print L1 and L2 values

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@551 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-16 16:34:38 +00:00
parent 149385d292
commit 262ae30576
+45 -25
View File
@@ -115,15 +115,17 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(miniBatchSize, trainingSize));
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 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 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()));
MatrixXd penalty_weights = MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden());
double L1 = 0;
double L2 = 0;
if (m_params.m_doNormalizeData && !m_params.m_useVisibleGaussian)
{
probsLogistic(__batch);
@@ -138,15 +140,13 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
trainingSizeRemain -= toSlice;
batchRowIndex += toSlice;
size_t batchSize = batch.rows();
double mu_w = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
double mu_biasV = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
double mu_biasH = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
double learning_rate = m_params.m_muWeights/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());
for (epoch=0; epoch < numEpochs; epoch++)
{
onProgressChanged();
@@ -175,9 +175,9 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
}
// Update weights (positive phase)
dBiasV_curr = batch.colwise().sum();
dBiasH_curr = hid.colwise().sum();
dW_curr = batch.transpose() * hid;
grad_bias_v = batch.colwise().sum();
grad_bias_h = hid.colwise().sum();
grad_weight = batch.transpose() * hid;
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
{
@@ -225,12 +225,32 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
}
// Update weights (negative phase)
dBiasV_curr -= vis.colwise().sum();
dBiasH_curr -= hid.colwise().sum();
dW_curr -= vis.transpose() * hid;
grad_bias_v -= vis.colwise().sum();
grad_bias_h -= hid.colwise().sum();
grad_weight -= vis.transpose() * hid;
m_w.visibleBias() += mu_biasV*(m_params.m_momentum*dBiasV + (1-m_params.m_momentum)*dBiasV_curr);
dBiasV = dBiasV_curr;
for (int i=0; i < m_w.weights().rows(); i++)
{
for (int j=0; j < m_w.weights().cols(); j++)
{
if (m_w.weights()(i,j) >= 0)
{
penalty_weights(i,j) = m_params.m_weightDecay;
}
else
{
penalty_weights(i,j) = -m_params.m_weightDecay;
}
}
}
L1 = m_w.weights().array().abs().sum();
L2 = m_w.weights().array().square().sum();
momentum_bias_v = m_params.m_momentum*momentum_bias_v + grad_bias_v;
momentum_bias_h = m_params.m_momentum*momentum_bias_h + grad_bias_h;
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)
{
@@ -240,13 +260,11 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
}
else
{
m_w.hiddenBias() += mu_biasH*(m_params.m_momentum*dBiasH + (1-m_params.m_momentum)*dBiasH_curr);
m_w.hiddenBias() += learning_rate*momentum_bias_h;
}
dBiasH = dBiasH_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()*m_w.weights().array().abs().sum());
dW = dW_curr;
m_w.weights() += learning_rate*momentum_weights;
if (m_params.m_sigmaDecay > 0)
{
if (m_variableSigma[0] > sigmaMin)
@@ -262,6 +280,8 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
diffErr.array() *= diffErr.array();
double err = diffErr.colwise().sum().sum();
cout << "error (per mini batch) = " << err << endl;
cout << "L1 = " << L1 << endl;
cout << "L2 = " << L2 << endl;
} // number of mini batches
updateHiddenBatch();