From 09963e2b6316c95abc590557954db4bb4c1f6f55 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 21 Oct 2019 22:13:49 +0000 Subject: [PATCH] Improved status git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@562 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Rbm.cpp | 46 ++++++++++++++++++++++------------------------ source/Rbm.hpp | 40 +++++++++++++++++++++++++--------------- source/main.cpp | 26 ++++++++++++++++++++++++-- 3 files changed, 71 insertions(+), 41 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 3f8e7fd..9698955 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -11,7 +11,6 @@ * Created on 21. Oktober 2019, 21:28 */ -#include #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 } diff --git a/source/Rbm.hpp b/source/Rbm.hpp index ce33454..bbf1fab 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -16,23 +16,11 @@ #include #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); diff --git a/source/main.cpp b/source/main.cpp index 423e15c..128c865 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -1,8 +1,28 @@ +#include #include #include #include #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);