diff --git a/source/Layer.cpp b/source/Layer.cpp index 1812a1b..d478b59 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -105,12 +105,9 @@ arma::mat Layer::upPass(const arma::mat& v) arma::mat h = gibbsPass(r); if (next) { - next->upPass(h); - } - else - { - return h; + return next->upPass(h); } + return h; } arma::mat Layer::upDownPass(const arma::mat& v) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 63447f3..b39e9cc 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -40,6 +40,50 @@ Rbm::~Rbm() Noise_Free(&m_noise); } +size_t Rbm::numHidden() const +{ + return m_bhv.size(); +} + +size_t Rbm::numVisible() const +{ + return m_bv.size(); +} + +Rbm::Params& Rbm::params() +{ + return m_params; +} + +arma::mat Rbm::rms_error(arma::mat diffErr) +{ + arma::mat diffErr_squared = diffErr % diffErr; + return arma::sum(diffErr_squared, 1) * 1.0 / diffErr_squared.n_cols; +} + +double Rbm::rms_error_accu(arma::mat diffErr) +{ + arma::mat diffErr_squared = diffErr % diffErr; + return arma::accu(diffErr_squared) / diffErr_squared.n_elem; +} + +arma::mat Rbm::toHiddenProbs(const arma::mat& visible) const +{ + return Rbm::prob(v_to_h(visible)); +} + +arma::mat Rbm::toVisibleProbs(const arma::mat& hidden) const +{ + return Rbm::prob(h_to_v(hidden)); +} + +void Rbm::weightsAssign(const arma::mat& w, const arma::mat& bhv, const arma::mat& bv) +{ + m_whv.submat(0, 0, w.n_rows - 1, w.n_cols - 1) = w; + m_bhv.submat(0, 0, bhv.n_rows - 1, bhv.n_cols - 1) = bhv; + m_bv.submat(0, 0, bv.n_rows - 1, bv.n_cols - 1) = bv; +} + void Rbm::weightsInit(double stddev, double mu) { uniform(m_whv, stddev, mu); diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 7024141..8b0ceff 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -23,7 +23,7 @@ class Rbm { public: - + struct Params { Params() @@ -72,7 +72,7 @@ public: miniBatchSize = params.get("miniBatchSize", miniBatchSize).asUInt(); numEpochs = params.get("numEpochs", numEpochs).asUInt(); } - + double weightDecay; double learningRate; double momentum; @@ -110,7 +110,7 @@ public: virtual ~IListener() {} virtual bool onProgress(Rbm *pRbm, const Status &status) { - return true; + return true; } }; @@ -118,73 +118,44 @@ public: Rbm(const Rbm& orig); virtual ~Rbm(); + Params& params(); + void weightsInit(double stddev, double mu=0.0); - void weightsAssign(const arma::mat &w, const arma::mat &bhv, const arma::mat &bv) - { - m_whv.submat(0, 0, w.n_rows-1, w.n_cols-1) = w; - m_bhv.submat(0, 0, bhv.n_rows-1, bhv.n_cols-1) = bhv; - m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv; - } - + void weightsAssign(const arma::mat &w, const arma::mat &bhv, const arma::mat &bv); + 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; - + Json::Value toJson() const; void fromJson(Json::Value params); - - Params& params() - { - return m_params; - } - static arma::mat prob(arma::mat const &src); - - 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)); - } - - size_t numHidden() const - { - return m_bhv.size(); - } + arma::mat toHiddenProbs(const arma::mat &visible) const; + arma::mat toVisibleProbs(const arma::mat &hidden) const; + + size_t numHidden() const; + + size_t numVisible() const; - size_t numVisible() const - { - return m_bv.size(); - } - arma::mat v_to_h(const arma::mat &visible) const; arma::mat h_to_v(const arma::mat &hidden) const; + static arma::mat prob(arma::mat const &src); - static double rms_error_accu(arma::mat diffErr) - { - arma::mat diffErr_squared = diffErr % diffErr; - return arma::accu(diffErr_squared)/diffErr_squared.n_elem; - } + static double rms_error_accu(arma::mat diffErr); + static arma::mat rms_error(arma::mat diffErr); - static arma::mat rms_error(arma::mat diffErr) - { - arma::mat diffErr_squared = diffErr % diffErr; - return arma::sum(diffErr_squared, 1) * 1.0/diffErr_squared.n_cols; - } void gibbs_hv(arma::mat &h_states, arma::mat &v_states); void gibbs_vh(arma::mat &v_probs, arma::mat &h_probs); - + 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: arma::mat m_bhv; arma::mat m_bv;