- Rbm: cleaned up, learning_rate and weight_decay normalized to miniBatchSizeActual
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@590 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+4
-5
@@ -63,11 +63,10 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
|
||||
size_t epoch;
|
||||
size_t gibbs;
|
||||
|
||||
size_t trainingSize = batch.n_rows;
|
||||
status.trainingSizeRemain = trainingSize;
|
||||
status.trainingSizeRemain = batch.n_rows;
|
||||
size_t batchRowIndex = 0;
|
||||
|
||||
double dProgress = 1.0/trainingSize;
|
||||
double dProgress = 1.0/status.trainingSizeRemain;
|
||||
|
||||
arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows));
|
||||
arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols));
|
||||
@@ -88,8 +87,8 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
|
||||
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
||||
status.trainingSizeRemain -= miniBatchSizeActual;
|
||||
batchRowIndex += miniBatchSizeActual;
|
||||
double learning_rate = m_params.learningRate/std::min(miniBatchSizeActual, trainingSize);
|
||||
double weight_decay = m_params.weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
||||
double learning_rate = m_params.learningRate/miniBatchSizeActual;
|
||||
double weight_decay = m_params.weightDecay/miniBatchSizeActual;
|
||||
|
||||
arma::mat vis_state(miniBatchSizeActual, m_w.n_rows);
|
||||
arma::mat vis_probs(miniBatchSizeActual, m_w.n_rows);
|
||||
|
||||
Reference in New Issue
Block a user