Improved status
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@562 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+22
-24
@@ -11,7 +11,6 @@
|
||||
* Created on 21. Oktober 2019, 21:28
|
||||
*/
|
||||
|
||||
#include <streambuf>
|
||||
#include "Rbm.hpp"
|
||||
#include "noise.h"
|
||||
|
||||
@@ -34,9 +33,8 @@ Rbm::~Rbm()
|
||||
{
|
||||
}
|
||||
|
||||
void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize, const Params& params, IRbmListener* pListener)
|
||||
void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize, const Params& params, IListener* pListener)
|
||||
{
|
||||
size_t i;
|
||||
size_t epoch;
|
||||
size_t gibbs;
|
||||
|
||||
@@ -53,13 +51,13 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
arma::mat momentum_bias_v(arma::zeros(1, m_w.n_rows));
|
||||
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);
|
||||
double L1 = 0;
|
||||
double L2 = 0;
|
||||
|
||||
double progress = 0;
|
||||
Status status;
|
||||
|
||||
status.progress = 0;
|
||||
while (trainingSizeRemain)
|
||||
{
|
||||
std::cout << "trainingSizeRemain: " << trainingSizeRemain << std::endl;
|
||||
status.trainingSizeRemain = trainingSizeRemain;
|
||||
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
|
||||
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+toSlice-1);
|
||||
trainingSizeRemain -= toSlice;
|
||||
@@ -75,14 +73,7 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
|
||||
for (epoch=0; epoch < numEpochs; epoch++)
|
||||
{
|
||||
if (pListener)
|
||||
{
|
||||
if(!pListener->onProgress())
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Create hidden layer base on training data
|
||||
if (params.m_doSampleBatch)
|
||||
{
|
||||
@@ -156,26 +147,33 @@ void Rbm::train(const arma::mat& batch, size_t numEpochs, size_t miniBatchSize,
|
||||
}
|
||||
}
|
||||
|
||||
L1 = accu(abs(m_w));
|
||||
L2 = accu(m_w % m_w);
|
||||
status.L1 = accu(abs(m_w));
|
||||
status.L2 = accu(m_w % m_w);
|
||||
momentum_bias_v = params.m_momentum*momentum_bias_v + grad_bias_v;
|
||||
momentum_bias_h = params.m_momentum*momentum_bias_h + grad_bias_h;
|
||||
momentum_weights = params.m_momentum*momentum_weights + grad_weight - L2*penalty_weights;
|
||||
momentum_weights = params.m_momentum*momentum_weights + grad_weight - status.L2*penalty_weights;
|
||||
|
||||
m_bv += learning_rate*momentum_bias_v;
|
||||
m_bh += learning_rate*momentum_bias_h;
|
||||
m_w += learning_rate*momentum_weights;
|
||||
|
||||
progress += dProgress;
|
||||
|
||||
status.progress += dProgress;
|
||||
|
||||
} // Number of epochs
|
||||
|
||||
status.epoch = epoch;
|
||||
arma::mat diffErr = miniBatch - vis_probs;
|
||||
arma::mat diffErr_squared = diffErr % diffErr;
|
||||
double err = accu(diffErr_squared);
|
||||
std::cout << "error (per mini batch) = " << err << std::endl;
|
||||
std::cout << "L1 = " << L1 << std::endl;
|
||||
std::cout << "L2 = " << L2 << std::endl;
|
||||
status.err = accu(diffErr_squared);
|
||||
|
||||
if (pListener)
|
||||
{
|
||||
if(!pListener->onProgress(status))
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
} // number of mini batches
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user