From e240eae073dcfaa9c59734669f5467a15b336f1b Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 10 Jan 2022 10:03:43 +0000 Subject: [PATCH] - moved context awareness from Rbm to Layer (final) git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@770 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Layer.cpp | 5 +++-- source/Layer.hpp | 26 +++++++++++++++++++++++--- source/Rbm.cpp | 42 +++++++++--------------------------------- source/Rbm.hpp | 21 +++------------------ 4 files changed, 38 insertions(+), 56 deletions(-) diff --git a/source/Layer.cpp b/source/Layer.cpp index 5660cb9..3ca066f 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -15,13 +15,14 @@ using namespace std; Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext) -: Rbm(numVisibleX*numVisibleY, numHidden, numContext) +: Rbm(numVisibleX*numVisibleY+numContext, numHidden) , next(nullptr) , prev(nullptr) , m_name(name) , m_id(id) , m_numVisibleX(numVisibleX) , m_numVisibleY(numVisibleY) +, m_numContext(numContext) { cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl; m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat"); @@ -164,7 +165,7 @@ Json::Value Layer::toJson() const layer["numVisibleX"] = (int)m_numVisibleX; layer["numVisibleY"] = (int)m_numVisibleY; layer["numHidden"] = (int)whv().n_cols; - layer["numContext"] = numContext(); + layer["numContext"] = m_numContext; layer["rbm"] = Rbm::toJson(); return layer; diff --git a/source/Layer.hpp b/source/Layer.hpp index a5998dc..69ffe90 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -38,10 +38,10 @@ public: { if (batch.n_rows > 0) { - if (numContext()) + if (m_numContext) { - arma::mat c_states = arma::zeros(batch.n_rows, numContext()); - arma::mat ctx = arma::zeros(1, numContext()); + arma::mat c_states = arma::zeros(batch.n_rows, m_numContext); + arma::mat ctx = arma::zeros(1, m_numContext); for (int i=1; i < batch.n_rows; i++) { arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx))); @@ -109,6 +109,25 @@ public: return pLayer; } + arma::mat vc_to_v(const arma::mat &vc) const + { + return arma::reshape(vc, 1, numVisible() - m_numContext); + } + + arma::mat vc_to_c(const arma::mat &vc) const + { + if (m_numContext == 0) + { + return arma::mat(1,0); + } + return vc.submat(0, numVisible() - m_numContext, 0, numVisible() - 1); + } + + arma::mat to_vc(const arma::mat &v, const arma::mat &c) const + { + return arma::join_rows(v, c); + } + private: std::string m_name; @@ -116,6 +135,7 @@ private: size_t m_id; size_t m_numVisibleX; size_t m_numVisibleY; + size_t m_numContext; }; diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 948ea15..3a7f676 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -16,20 +16,14 @@ #define RBM_TRAIN_FLAT 0 -Rbm::Rbm(size_t numVisible, size_t numHidden, size_t numContext) +Rbm::Rbm(size_t numVisible, size_t numHidden) : m_params() -, m_whv(numVisible+numContext, numHidden) +, m_whv(numVisible, numHidden) , m_bhv(1, numHidden) -, m_bv(1, numVisible+numContext) -, m_ctx() +, m_bv(1, numVisible) { assert(numVisible > 0); assert(numHidden > 0); - if (numContext) - { - assert(numContext == numHidden); - m_ctx.resize(1, numContext); - } Noise_Init(&m_noise, 0x32727155); } @@ -38,7 +32,6 @@ Rbm::Rbm(const Rbm& orig) , m_whv(orig.m_whv) , m_bhv(orig.m_bhv) , m_bv(orig.m_bv) -, m_ctx(orig.m_ctx) { } @@ -52,7 +45,6 @@ void Rbm::weightsInit(double stddev, double mu) uniform(m_whv, stddev, mu); uniform(m_bhv, stddev, mu); uniform(m_bv, stddev, mu); - uniform(m_ctx, stddev, mu); } void Rbm::fromJson(Json::Value rbm) @@ -95,7 +87,7 @@ void Rbm::gibbs(arma::mat &hv_probs, arma::mat &v_probs) } } -void Rbm::weightUpdate(arma::mat const &v_states, arma::mat &dw, arma::mat &dbhv, arma::mat &dbv) +void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::mat &dbhv, arma::mat &dbv) { arma::mat v_probs(v_states); arma::mat h_states = v_to_h(v_states); @@ -139,7 +131,6 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) arma::mat momentum_bias_v(arma::zeros(1, m_bv.n_cols)); arma::mat momentum_bias_hv(arma::zeros(1, m_bhv.n_cols)); arma::mat penalty_weights = arma::zeros(m_whv.n_rows, m_whv.n_cols); - arma::mat ctx = arma::zeros(1, numContext()); int trainingSizeRemain = batch.n_rows; @@ -165,8 +156,10 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) for (int epoch=0; epoch < m_params.numEpochs; epoch++) { - weightUpdate(v_states, grad_weight_hv, grad_bias_hv, grad_bias_v); + // Contrastive divergence learning: calculate gradients + contrastiveDivergence(v_states, grad_weight_hv, grad_bias_hv, grad_bias_v); + // Adjust weight and biases penalty_weights = weight_decay*arma::sign(m_whv); status.L1 = accu(abs(m_whv)); @@ -182,6 +175,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) progress += dProgress*miniBatchSizeActual; status.progress = (int)(progress + 0.5); + // Update status if (status.progress != lastProgress) { lastProgress = status.progress; @@ -204,6 +198,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) } // number of mini batches + // Update final status 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; @@ -234,25 +229,6 @@ arma::mat Rbm::sample(const arma::mat &src) return dst; } -arma::mat Rbm::vc_to_v(const arma::mat &vc) const -{ - return arma::reshape(vc, 1, numVisible() - numContext()); -} - -arma::mat Rbm::vc_to_c(const arma::mat &vc) const -{ - if (numContext() == 0) - { - return arma::mat(1,0); - } - return vc.submat(0, numVisible() - numContext(), 0, numVisible() - 1); -} - -arma::mat Rbm::to_vc(const arma::mat &v, const arma::mat &c) const -{ - return arma::join_rows(v, c); -} - arma::mat Rbm::v_to_h(const arma::mat &visible) const { return visible * m_whv + arma::repmat(m_bhv, visible.n_rows, 1); diff --git a/source/Rbm.hpp b/source/Rbm.hpp index c5389b7..163f971 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -114,7 +114,7 @@ public: } }; - Rbm(size_t numVisible, size_t numHidden, size_t numContext=0); + Rbm(size_t numVisible, size_t numHidden); Rbm(const Rbm& orig); virtual ~Rbm(); @@ -149,19 +149,9 @@ public: arma::mat toVisibleProbs(const arma::mat &hidden) const { - return arma::reshape(Rbm::prob(h_to_v(hidden)), 1, numVisible() - numContext()); + return Rbm::prob(h_to_v(hidden)); } - arma::mat toContextProbs(const arma::mat &hidden) const - { - return Rbm::prob(h_to_v(hidden)).submat(0, numVisible() - numContext(), 0, numVisible() - 1); - } - - size_t numContext() const - { - return m_ctx.size(); - } - size_t numHidden() const { return m_bhv.size(); @@ -172,17 +162,13 @@ public: return m_bv.size(); } - arma::mat vc_to_v(const arma::mat &vc) const; - arma::mat vc_to_c(const arma::mat &vc) const; - arma::mat to_vc(const arma::mat &v, const arma::mat &c) const; - arma::mat v_to_h(const arma::mat &visible) const; arma::mat h_to_v(const arma::mat &hidden) const; private: Params m_params; arma::mat sample(arma::mat const &src); - void weightUpdate(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); + 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); void gibbs(arma::mat &hv_states, arma::mat &v_states); @@ -193,7 +179,6 @@ protected: private: noise_gen_t m_noise; arma::mat m_whv; - arma::mat m_ctx; };