From aef3049856a5fc743a4e073e967d10c7c5188a2e Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 21 Jan 2022 07:46:40 +0000 Subject: [PATCH] - refactored common functions into matutils git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@855 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/AStack.cpp | 5 +++- source/Rbm.cpp | 51 ++------------------------------ source/Rbm.hpp | 9 ++---- source/matutils.hpp | 72 +++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 80 insertions(+), 57 deletions(-) create mode 100644 source/matutils.hpp diff --git a/source/AStack.cpp b/source/AStack.cpp index 442d749..67a2992 100644 --- a/source/AStack.cpp +++ b/source/AStack.cpp @@ -15,7 +15,10 @@ #include "StackCreator.hpp" #include +#include "matutils.hpp" + using namespace std; +using namespace Matutils; const char *AStack::stackTypeStrings[NUM_STACKTYPES] = {"None", "Deep", "Rnn"}; @@ -158,7 +161,7 @@ size_t AStack::loadTrainingBatch(const std::string &dir, bool doNormalize) { if (doNormalize) { - m_trainingBatch = Rbm::normalize(m_trainingBatch); + m_trainingBatch = normalize(m_trainingBatch); } std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; } diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 16201d2..6a68dbc 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -13,8 +13,9 @@ #include #include "Rbm.hpp" +#include "matutils.hpp" -#define RBM_TRAIN_FLAT 0 +using namespace Matutils; Rbm::Rbm(size_t numVisible, size_t numHidden) : m_params() @@ -272,27 +273,6 @@ arma::mat Rbm::prob(const arma::mat &src) return 1 / (1 + (arma::exp(-src))); } -arma::mat Rbm::sample(const arma::mat &src) -{ - arma::mat dst = src; - uniform(dst); - -#if 0 - for (size_t i=0; i < src.n_rows; i++) - { - for (size_t j=0; j < src.n_cols; j++) - { - dst(i, j) = src(i, j) >= dst(i, j); - } - } - return dst; -#else - arma::umat res = (dst < src); - return arma::conv_to::from(res); - -#endif -} - arma::mat Rbm::v_to_h(const arma::mat &visible) const { return visible * m_whv + arma::repmat(m_bhv, visible.n_rows, 1); @@ -303,33 +283,6 @@ arma::mat Rbm::h_to_v(const arma::mat &hidden) const return hidden * m_whv.t() + arma::repmat(m_bv, hidden.n_rows, 1); } -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 0 - for (size_t i=0; i < srcDst.n_rows; i++) - { - for (size_t j=0; j < srcDst.n_cols; j++) - { - srcDst(i, j) = stdDev*(Noise_Uniform(&m_noise) + mu - 0.5); - } - } -#else - srcDst = stdDev*(arma::randu(arma::size(srcDst)) + mu - 0.5); -#endif -} - const arma::mat& Rbm::whv() const { return m_whv; diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 1182186..1e4d8c6 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -124,7 +124,6 @@ public: void train(arma::mat const &batch, IListener *pListener=nullptr); - static arma::mat normalize(const arma::mat &hidden); const arma::mat& whv() const; const arma::mat& bv() const; const arma::mat& bh() const; @@ -148,18 +147,14 @@ public: static double rms_error_accu(arma::mat diffErr); static arma::mat rms_error(arma::mat diffErr); -private: - Params m_params; - arma::mat sample(arma::mat const &src); - void contrastiveDivergence(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); - void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5); - protected: + Params m_params; arma::mat m_bhv; arma::mat m_bv; private: arma::mat m_whv; + void contrastiveDivergence(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); }; diff --git a/source/matutils.hpp b/source/matutils.hpp new file mode 100644 index 0000000..f7468c5 --- /dev/null +++ b/source/matutils.hpp @@ -0,0 +1,72 @@ +/* + * To change this license header, choose License Headers in Project Properties. + * To change this template file, choose Tools | Templates + * and open the template in the editor. + */ + +/* + * File: matutils.hpp + * Author: jens + * + * Created on 21. Januar 2022, 08:28 + */ + +#ifndef MATUTILS_HPP +#define MATUTILS_HPP + +#include + +namespace Matutils +{ + inline void uniform(arma::mat& srcDst, double stdDev=1.0, double mu=0.5) + { + #if 0 + for (size_t i=0; i < srcDst.n_rows; i++) + { + for (size_t j=0; j < srcDst.n_cols; j++) + { + srcDst(i, j) = stdDev*(Noise_Uniform(&m_noise) + mu - 0.5); + } + } + #else + srcDst = stdDev*(arma::randu(arma::size(srcDst)) + mu - 0.5); + #endif + } + + inline arma::mat sample(const arma::mat &src) + { + arma::mat dst = src; + uniform(dst); + + #if 0 + for (size_t i=0; i < src.n_rows; i++) + { + for (size_t j=0; j < src.n_cols; j++) + { + dst(i, j) = src(i, j) >= dst(i, j); + } + } + return dst; + #else + arma::umat res = (dst < src); + return arma::conv_to::from(res); + + #endif + } + + inline arma::mat 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; + } + +} + +#endif /* MATUTILS_HPP */ +