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
|
||||
}
|
||||
|
||||
|
||||
+25
-15
@@ -16,23 +16,11 @@
|
||||
|
||||
#include <armadillo>
|
||||
#include "noise.h"
|
||||
class IRbmListener
|
||||
{
|
||||
public:
|
||||
|
||||
virtual ~IRbmListener()
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
bool onProgress()
|
||||
{
|
||||
return true;
|
||||
}
|
||||
};
|
||||
class Rbm
|
||||
{
|
||||
public:
|
||||
|
||||
struct Params
|
||||
{
|
||||
Params()
|
||||
@@ -45,7 +33,7 @@ public:
|
||||
, m_numGibbs(1)
|
||||
{
|
||||
}
|
||||
|
||||
|
||||
double m_weightDecay;
|
||||
double m_learningRate;
|
||||
double m_momentum;
|
||||
@@ -55,11 +43,33 @@ public:
|
||||
size_t m_numGibbs;
|
||||
};
|
||||
|
||||
struct Status
|
||||
{
|
||||
size_t epoch;
|
||||
size_t trainingSizeRemain;
|
||||
double progress;
|
||||
double err;
|
||||
double L1;
|
||||
double L2;
|
||||
};
|
||||
|
||||
class IListener
|
||||
{
|
||||
public:
|
||||
|
||||
IListener() {}
|
||||
virtual ~IListener() {}
|
||||
virtual bool onProgress(const Status &status)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
Rbm(size_t numHidden, size_t numVisible);
|
||||
Rbm(const Rbm& orig);
|
||||
virtual ~Rbm();
|
||||
|
||||
void train(arma::mat const &batch, size_t numEpochs, size_t sizeMiniBatch, Params const ¶ms, IRbmListener *pListener);
|
||||
void train(arma::mat const &batch, size_t numEpochs, size_t sizeMiniBatch, Params const ¶ms, IListener *pListener);
|
||||
arma::mat sample(arma::mat const &src);
|
||||
static arma::mat probsLogistic(arma::mat const &src);
|
||||
arma::mat toHidden(const arma::mat &v);
|
||||
|
||||
+24
-2
@@ -1,8 +1,28 @@
|
||||
#include <streambuf>
|
||||
#include <cstdio>
|
||||
#include <cmath>
|
||||
#include <armadillo>
|
||||
#include "Rbm.hpp"
|
||||
|
||||
class RbmListener : public Rbm::IListener
|
||||
{
|
||||
public:
|
||||
RbmListener() {}
|
||||
virtual ~RbmListener() {}
|
||||
|
||||
bool onProgress(const Rbm::Status &status)
|
||||
{
|
||||
std::cout << "Progress : " << 100*status.progress << " %" << std::endl;
|
||||
std::cout << "epoch : " << status.epoch << std::endl;
|
||||
std::cout << "trainingSizeRemain: " << status.trainingSizeRemain << std::endl;
|
||||
std::cout << "error (per mini batch) = " << status.err << std::endl;
|
||||
std::cout << "L1 = " << status.L1 << std::endl;
|
||||
std::cout << "L2 = " << status.L2 << std::endl;
|
||||
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
int main()
|
||||
{
|
||||
printf("Hallo, Welt!\n");
|
||||
@@ -31,10 +51,12 @@ int main()
|
||||
printf("Z.n_elem = %d\n", (int)Z.n_elem);
|
||||
// +------>
|
||||
|
||||
|
||||
RbmListener statusDisplay;
|
||||
Rbm::Params params;
|
||||
Rbm rbm(28*28, 64);
|
||||
arma::mat batch = arma::randu(1000, 28*28);
|
||||
rbm.train(batch, 1000, 100, params, nullptr);
|
||||
arma::mat batch = arma::randu(100, 28*28);
|
||||
rbm.train(batch, 100, 10, params, &statusDisplay);
|
||||
|
||||
// arma::mat v = arma::randu(28*28, 1);
|
||||
// arma::mat h = rbm.toHidden(v);
|
||||
|
||||
Reference in New Issue
Block a user