- improved progressIndicator
- fixed Stack::trainingData() - fixed test::main git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@638 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+14
-16
@@ -63,11 +63,10 @@ Json::Value Rbm::toJson() const
|
||||
void Rbm::train(const arma::mat& batch, IListener* pListener)
|
||||
{
|
||||
Status status;
|
||||
size_t epoch;
|
||||
size_t gibbs;
|
||||
|
||||
status.trainingSizeRemain = batch.n_rows;
|
||||
size_t batchRowIndex = 0;
|
||||
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
|
||||
double progress = 0;
|
||||
int lastProgress = -100;
|
||||
int batchRowIndex = 0;
|
||||
|
||||
arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows));
|
||||
arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols));
|
||||
@@ -77,12 +76,14 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
||||
arma::mat momentum_bias_h(arma::zeros(1, m_w.n_cols));
|
||||
arma::mat penalty_weights = arma::zeros(m_w.n_rows, m_w.n_cols);
|
||||
|
||||
int trainingSizeRemain = batch.n_rows;
|
||||
|
||||
bool shouldAbort = false;
|
||||
while (status.trainingSizeRemain && !shouldAbort)
|
||||
while (trainingSizeRemain && !shouldAbort)
|
||||
{
|
||||
size_t miniBatchSizeActual = std::min(m_params.miniBatchSize, status.trainingSizeRemain);
|
||||
int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain);
|
||||
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
||||
status.trainingSizeRemain -= miniBatchSizeActual;
|
||||
trainingSizeRemain -= miniBatchSizeActual;
|
||||
batchRowIndex += miniBatchSizeActual;
|
||||
double learning_rate = m_params.learningRate/miniBatchSizeActual;
|
||||
double weight_decay = m_params.weightDecay/miniBatchSizeActual;
|
||||
@@ -92,13 +93,11 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
||||
arma::mat hid_state(miniBatchSizeActual, m_w.n_cols);
|
||||
arma::mat hid_probs(miniBatchSizeActual, m_w.n_cols);
|
||||
|
||||
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
|
||||
double lastProgress = -100.0;
|
||||
|
||||
for (epoch=0; epoch < m_params.numEpochs; epoch++)
|
||||
for (int epoch=0; epoch < m_params.numEpochs; epoch++)
|
||||
{
|
||||
|
||||
if ((status.progress - lastProgress) >= 1.00)
|
||||
status.progress = (int)(progress + 0.5);
|
||||
if (status.progress != lastProgress)
|
||||
{
|
||||
lastProgress = status.progress;
|
||||
if (pListener)
|
||||
@@ -139,7 +138,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
||||
grad_bias_v = sum(vis_state, 0);
|
||||
grad_bias_h = sum(hid_state, 0);
|
||||
|
||||
for (gibbs=0; gibbs < m_params.numGibbs; gibbs++)
|
||||
for (int gibbs=0; gibbs < m_params.numGibbs; gibbs++)
|
||||
{
|
||||
// Create visible reconstruction (a fantasy...) given hid
|
||||
if (m_params.gibbsDoSampleHidden)
|
||||
@@ -181,11 +180,10 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
|
||||
m_bh += learning_rate*momentum_bias_h;
|
||||
m_w += learning_rate*momentum_weights;
|
||||
|
||||
status.progress += dProgress*miniBatchSizeActual;
|
||||
progress += dProgress*miniBatchSizeActual;
|
||||
|
||||
} // Number of epochs
|
||||
|
||||
status.epoch = epoch;
|
||||
arma::mat diffErr = miniBatch - vis_probs;
|
||||
arma::mat diffErr_squared = diffErr % diffErr;
|
||||
status.err = accu(diffErr_squared)/diffErr_squared.n_elem;
|
||||
|
||||
Reference in New Issue
Block a user