- 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:
+42
-22
@@ -115,14 +115,16 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
|||||||
|
|
||||||
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(miniBatchSize, trainingSize));
|
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(miniBatchSize, trainingSize));
|
||||||
|
|
||||||
MatrixXd dBiasV_curr(MatrixXd::Zero(1, m_w.getNumVisible()));
|
MatrixXd grad_bias_v(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||||
MatrixXd dBiasH_curr(MatrixXd::Zero(1, m_w.getNumHidden()));
|
MatrixXd grad_bias_h(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||||
MatrixXd dW_curr(MatrixXd::Zero(m_w.getNumVisible(), m_w.getNumHidden()));
|
MatrixXd grad_weight(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 __batch = m_batch_normalized;
|
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)
|
if (m_params.m_doNormalizeData && !m_params.m_useVisibleGaussian)
|
||||||
{
|
{
|
||||||
@@ -138,9 +140,7 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
|||||||
trainingSizeRemain -= toSlice;
|
trainingSizeRemain -= toSlice;
|
||||||
batchRowIndex += toSlice;
|
batchRowIndex += toSlice;
|
||||||
size_t batchSize = batch.rows();
|
size_t batchSize = batch.rows();
|
||||||
double mu_w = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
double learning_rate = 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);
|
|
||||||
|
|
||||||
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
||||||
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
||||||
@@ -175,9 +175,9 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Update weights (positive phase)
|
// Update weights (positive phase)
|
||||||
dBiasV_curr = batch.colwise().sum();
|
grad_bias_v = batch.colwise().sum();
|
||||||
dBiasH_curr = hid.colwise().sum();
|
grad_bias_h = hid.colwise().sum();
|
||||||
dW_curr = batch.transpose() * hid;
|
grad_weight = batch.transpose() * hid;
|
||||||
|
|
||||||
for (gibbs=0; gibbs < m_params.m_numGibbs; gibbs++)
|
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)
|
// Update weights (negative phase)
|
||||||
dBiasV_curr -= vis.colwise().sum();
|
grad_bias_v -= vis.colwise().sum();
|
||||||
dBiasH_curr -= hid.colwise().sum();
|
grad_bias_h -= hid.colwise().sum();
|
||||||
dW_curr -= vis.transpose() * hid;
|
grad_weight -= vis.transpose() * hid;
|
||||||
|
|
||||||
m_w.visibleBias() += mu_biasV*(m_params.m_momentum*dBiasV + (1-m_params.m_momentum)*dBiasV_curr);
|
for (int i=0; i < m_w.weights().rows(); i++)
|
||||||
dBiasV = dBiasV_curr;
|
{
|
||||||
|
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)
|
if (m_params.m_doSparse)
|
||||||
{
|
{
|
||||||
@@ -240,12 +260,10 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
|||||||
}
|
}
|
||||||
else
|
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());
|
m_w.weights() += learning_rate*momentum_weights;
|
||||||
dW = dW_curr;
|
|
||||||
|
|
||||||
if (m_params.m_sigmaDecay > 0)
|
if (m_params.m_sigmaDecay > 0)
|
||||||
{
|
{
|
||||||
@@ -262,6 +280,8 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
|||||||
diffErr.array() *= diffErr.array();
|
diffErr.array() *= diffErr.array();
|
||||||
double err = diffErr.colwise().sum().sum();
|
double err = diffErr.colwise().sum().sum();
|
||||||
cout << "error (per mini batch) = " << err << endl;
|
cout << "error (per mini batch) = " << err << endl;
|
||||||
|
cout << "L1 = " << L1 << endl;
|
||||||
|
cout << "L2 = " << L2 << endl;
|
||||||
} // number of mini batches
|
} // number of mini batches
|
||||||
|
|
||||||
updateHiddenBatch();
|
updateHiddenBatch();
|
||||||
|
|||||||
Reference in New Issue
Block a user