- 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:
2019-11-08 06:38:09 +00:00
parent 215cfd11fa
commit 9ea6e48fff
5 changed files with 37 additions and 44 deletions
+14 -16
View File
@@ -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;