From f610d2930efb1014b29adb9ceb27989f0f625120 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 11 Nov 2019 21:08:48 +0000 Subject: [PATCH] - add Rbm::normalize() - normalize of training data git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@660 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/MainComponent.hpp | 2 +- source/Rbm.cpp | 12 ++++++++++++ source/Rbm.hpp | 1 + source/Stack.cpp | 7 ++++++- source/Stack.hpp | 2 +- 5 files changed, 21 insertions(+), 3 deletions(-) diff --git a/source/MainComponent.hpp b/source/MainComponent.hpp index 7c416f9..fe53d9f 100644 --- a/source/MainComponent.hpp +++ b/source/MainComponent.hpp @@ -100,7 +100,7 @@ private: void loadTraining() { - m_stack->loadTraining(); + m_stack->loadTraining(rbmNormalizeDataToggleButton->getToggleState()); patterSlider->setRange(0, m_stack->numTraining()-1, 1); } diff --git a/source/Rbm.cpp b/source/Rbm.cpp index f113995..5eae51d 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -243,6 +243,18 @@ arma::mat Rbm::toVisibleProbs(const arma::mat &hidden) const return probsLogistic(toVisibleState(hidden)); } +arma::mat Rbm::normalize(const arma::mat& src) +{ + double mean = arma::accu(src)/src.n_elem; + arma::mat x = src - mean; + arma::mat x2 = x % x; + double stddev = sqrt(arma::accu(x2)/x2.n_elem); + + std::cout << "mean" << " : " << std::endl << mean << std::endl; + std::cout << "stddev" << ": " << std::endl << stddev << std::endl; + return x/stddev; +} + void Rbm::uniform(arma::mat& srcDst, double stdDev, double mu) { #if 1 diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 22fbea3..19510aa 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -125,6 +125,7 @@ public: arma::mat toVisibleState(const arma::mat &hidden) const; arma::mat toHiddenProbs(const arma::mat &visible) const; arma::mat toVisibleProbs(const arma::mat &hidden) const; + static arma::mat normalize(const arma::mat &hidden); const arma::mat& w() const; const arma::mat& bv() const; const arma::mat& bh() const; diff --git a/source/Stack.cpp b/source/Stack.cpp index 61b0afe..9b59f70 100644 --- a/source/Stack.cpp +++ b/source/Stack.cpp @@ -222,7 +222,7 @@ arma::mat Stack::trainingData(Layer* pThatLayer) return thisBatch; } -size_t Stack::loadTraining() +size_t Stack::loadTraining(bool doNormalize) { uint32_t numTraining = 0; uint32_t numVisible = 0; @@ -263,6 +263,11 @@ size_t Stack::loadTraining() } fclose(pFile); std::cout << "Loaded " << numTraining << " training samples\n"; + + if (doNormalize) + { + m_trainingData = Rbm::normalize(m_trainingData); + } return numTraining; } diff --git a/source/Stack.hpp b/source/Stack.hpp index 4e59d8c..4eca74e 100644 --- a/source/Stack.hpp +++ b/source/Stack.hpp @@ -54,7 +54,7 @@ public: size_t numTraining(); void addTraining(const arma::mat &toAdd); void delTraining(int index); - size_t loadTraining(); + size_t loadTraining(bool doNormalize=false); size_t saveTraining(); arma::mat& trainingData(); arma::mat trainingData(Layer *pLayer);