From 637b237966a0aef58882353616a692e65b6df482 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sat, 8 Jan 2022 10:41:41 +0000 Subject: [PATCH] - refactored git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@751 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- Makefile | 4 ++-- source/Layer.hpp | 10 ++++++++++ source/Rbm.cpp | 42 +++++++++++++++++++++--------------------- source/Rbm.hpp | 11 ++++++----- 4 files changed, 39 insertions(+), 28 deletions(-) diff --git a/Makefile b/Makefile index d51aa52..92352ac 100644 --- a/Makefile +++ b/Makefile @@ -12,7 +12,7 @@ TEST_SRCS := source/main.cpp source/Rbm.cpp source/Layer.cpp source/Stack.cpp so CXXFLAGS += -std=c++11 -L juce/build/${CONFIG} CXXFLAGS_debug := ${CXXFLAGS} -O0 -g -CXXFLAGS_release := ${CXXFLAGS} -O2 +CXXFLAGS_release := ${CXXFLAGS} -O2 -g LFLAGS := -L juce/build/${CONFIG} LFLAGS_release := ${LFLAGS} @@ -20,7 +20,7 @@ LFLAGS_debug := ${LFLAGS} CFLAGS := CFLAGS_debug := ${CFLAGS} -O0 -g -CFLAGS_release := ${CFLAGS} -O2 +CFLAGS_release := ${CFLAGS} -O2 -g INCLUDES := -I${JUCE_ROOT} -Ijuce LIBS := -lstdc++ -lm -larmadillo -ljsoncpp -ljuce -lfreetype -lpthread -ldl -lrt -lX11 -lGL -lXinerama -lXext diff --git a/source/Layer.hpp b/source/Layer.hpp index b39a6f4..1950c60 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -67,6 +67,16 @@ public: return m_bh.n_elem; } + arma::mat toHiddenProbs(const arma::mat &visible) const + { + return Rbm::prob(v_to_h(visible)); + } + + arma::mat toVisibleProbs(const arma::mat &hidden) const + { + return Rbm::prob(h_to_v(hidden)); + } + arma::mat trainingData(arma::mat const &batch) { arma::mat thisBatch = batch; diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 1ee5f02..fbe7c55 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -73,7 +73,7 @@ Json::Value Rbm::toJson() const void Rbm::weightUpdate(arma::mat const &v_states, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) { - arma::mat h_probs = toHiddenProbs(v_states); + arma::mat h_probs = prob(v_to_h(v_states)); arma::mat v_probs(dbv.n_rows, dbv.n_cols); arma::mat h_states = h_probs; @@ -94,24 +94,24 @@ void Rbm::weightUpdate(arma::mat const &v_states, arma::mat &dw, arma::mat &dbh, // Create visible reconstruction (a fantasy...) given hid if (m_params.gibbsDoSampleHidden) { - v_probs = toVisibleProbs(sample(h_probs)); + v_probs = prob(h_to_v(sample(h_probs))); } else { - v_probs = toVisibleProbs(h_probs); + v_probs = prob(h_to_v(h_probs)); } // Create hidden representation given v if (m_params.gibbsDoSampleVisible) { - h_states = toHiddenState(sample(v_probs)); + h_states = v_to_h(sample(v_probs)); } else { - h_states = toHiddenState(v_probs); + h_states = v_to_h(v_probs); } - h_probs = probsLogistic(h_states); + h_probs = prob(h_states); } @@ -175,7 +175,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) { #if RBM_TRAIN_FLAT // Sample hidden - h_probs = probsLogistic(toHiddenState(v_states)); + h_probs = prob(v_to_h(v_states)); if (m_params.doRaoBlackwell) { hid_states = h_probs; @@ -195,23 +195,23 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) // Create visible reconstruction (a fantasy...) given hid if (m_params.gibbsDoSampleHidden) { - v_probs = toVisibleProbs(sample(h_probs)); + v_probs = prob(h_to_v(sample(h_probs))); } else { - v_probs = toVisibleProbs(h_probs); + v_probs = prob(h_to_v(h_probs)); } // Create hidden representation given v if (m_params.gibbsDoSampleVisible) { - hid_states = toHiddenState(sample(v_probs)); + hid_states = v_to_h(sample(v_probs)); } else { - hid_states = toHiddenState(v_probs); + hid_states = v_to_h(v_probs); } - h_probs = probsLogistic(hid_states); + h_probs = prob(hid_states); } @@ -243,7 +243,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) lastProgress = status.progress; // Calculate error - arma::mat diffErr = miniBatch - toVisibleProbs(toHiddenProbs(v_states)); + arma::mat diffErr = miniBatch - prob(h_to_v(prob(v_to_h(v_states)))); arma::mat diffErr_squared = diffErr % diffErr; status.err = accu(diffErr_squared)/diffErr_squared.n_elem; if (pListener) @@ -260,7 +260,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) } // number of mini batches - arma::mat diffErr = batch - toVisibleProbs(toHiddenProbs(batch)); + arma::mat diffErr = batch - prob(h_to_v(prob(v_to_h(batch)))); arma::mat diffErr_squared = diffErr % diffErr; status.err_total = accu(diffErr_squared)/diffErr_squared.n_elem; @@ -270,7 +270,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) } } -arma::mat Rbm::probsLogistic(const arma::mat &src) +arma::mat Rbm::prob(const arma::mat &src) { return 1 / (1 + (arma::exp(-src))); } @@ -290,24 +290,24 @@ arma::mat Rbm::sample(const arma::mat &src) return dst; } -arma::mat Rbm::toHiddenState(const arma::mat &visible) const +arma::mat Rbm::v_to_h(const arma::mat &visible) const { return visible * m_whv + arma::repmat(m_bh, visible.n_rows, 1); } -arma::mat Rbm::toVisibleState(const arma::mat &hidden) const +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::toHiddenProbs(const arma::mat &visible) const +arma::mat Rbm::c_to_h(const arma::mat &context) const { - return probsLogistic(toHiddenState(visible)); + return context * m_whc + arma::repmat(m_bh, context.n_rows, 1); } -arma::mat Rbm::toVisibleProbs(const arma::mat &hidden) const +arma::mat Rbm::h_to_c(const arma::mat &hidden) const { - return probsLogistic(toVisibleState(hidden)); + return hidden * m_whc.t() + arma::repmat(m_bc, hidden.n_rows, 1); } arma::mat Rbm::normalize(const arma::mat& src) diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 03687f1..66b24c2 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -121,10 +121,6 @@ public: void weightsInit(double stddev, double mu=0.0); void train(arma::mat const &batch, IListener *pListener=nullptr); - arma::mat toHiddenState(const arma::mat &visible) const; - 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; @@ -137,10 +133,15 @@ public: { return m_params; } + static arma::mat prob(arma::mat const &src); + arma::mat v_to_h(const arma::mat &visible) const; + arma::mat h_to_v(const arma::mat &hidden) const; + arma::mat c_to_h(const arma::mat &context) const; + arma::mat h_to_c(const arma::mat &hidden) const; + private: Params m_params; arma::mat sample(arma::mat const &src); - static arma::mat probsLogistic(arma::mat const &src); void weightUpdate(arma::mat const &v_state, arma::mat &dw, arma::mat &dbh, arma::mat &dbv); void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5);