diff --git a/Source/Rbm.cpp b/Source/Rbm.cpp index 4434dac..1937781 100644 --- a/Source/Rbm.cpp +++ b/Source/Rbm.cpp @@ -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();