Rbm:🚋 better status indication

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@584 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-26 06:26:02 +00:00
parent f50607ef74
commit 28bbff0b45
2 changed files with 20 additions and 8 deletions
+10 -8
View File
@@ -59,11 +59,12 @@ Json::Value Rbm::toJson() const
void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener)
{
Status status;
size_t epoch;
size_t gibbs;
size_t trainingSize = batch.n_rows;
size_t trainingSizeRemain = trainingSize;
status.trainingSizeRemain = trainingSize;
size_t batchRowIndex = 0;
double dProgress = 1.0/trainingSize;
@@ -76,15 +77,16 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
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);
Status status;
status.progress = 0;
while (trainingSizeRemain)
if (pListener)
{
status.trainingSizeRemain = trainingSizeRemain;
size_t miniBatchSizeActual = std::min(miniBatchSize, trainingSizeRemain);
pListener->onProgress(status);
}
while (status.trainingSizeRemain)
{
size_t miniBatchSizeActual = std::min(miniBatchSize, status.trainingSizeRemain);
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
trainingSizeRemain -= miniBatchSizeActual;
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);
+10
View File
@@ -81,6 +81,16 @@ public:
struct Status
{
Status()
: epoch(0)
, trainingSizeRemain(0)
, progress(0)
, err(-1.0)
, err_total(1-0)
, L1(-1.0)
, L2(-1.0)
{
}
size_t epoch;
size_t trainingSizeRemain;
double progress;