diff --git a/source/Layer.cpp b/source/Layer.cpp index eed9b57..e451269 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -89,6 +89,7 @@ bool Layer::loadWeights(const string &prjname) m_bhv(i) = v; } } + arma::mat weights(numVisible, numHidden); for (i=0; i < numVisible; i++) { for (j=0; j < numHidden; j++) @@ -97,10 +98,11 @@ bool Layer::loadWeights(const string &prjname) result = fscanf(pFile, "%f", &v); if (result > 0) { - m_whv(i, j) = v; + weights(i, j) = v; } } } + weightsAssign(weights); fclose(pFile); return true; @@ -136,11 +138,12 @@ bool Layer::saveWeights(const string &prjname) { fprintf(pFile, "%3.6f\n", m_bhv(i)); } + const arma::mat &_whv = whv(); for (i=0; i < numVisible; i++) { for (j=0; j < numHidden; j++) { - fprintf(pFile, "%3.6f ", m_whv(i,j)); + fprintf(pFile, "%3.6f ", _whv(i,j)); } fprintf(pFile, "\n"); } diff --git a/source/Layer.hpp b/source/Layer.hpp index 3d95386..1678d0a 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -67,16 +67,6 @@ public: return m_bhv.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 2419259..a39c7cc 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -18,12 +18,10 @@ Rbm::Rbm(size_t numVisible, size_t numHidden, size_t numContext) : m_params() -, m_whv(numVisible, numHidden) -, m_whc(numContext, numHidden) +, m_whv(numVisible+numContext, numHidden) , m_bhv(1, numHidden) -, m_bhc(1, numHidden) -, m_bv(1, numVisible) -, m_bc(1, numContext) +, m_bv(1, numVisible+numContext) +, m_ctx(1, numContext) { assert(numVisible > 0); assert(numHidden > 0); @@ -37,11 +35,9 @@ Rbm::Rbm(size_t numVisible, size_t numHidden, size_t numContext) Rbm::Rbm(const Rbm& orig) : m_params(orig.m_params) , m_whv(orig.m_whv) -, m_whc(orig.m_whc) , m_bhv(orig.m_bhv) -, m_bhc(orig.m_bhc) , m_bv(orig.m_bv) -, m_bc(orig.m_bc) +, m_ctx(orig.m_ctx) { } @@ -53,11 +49,9 @@ Rbm::~Rbm() void Rbm::weightsInit(double stddev, double mu) { uniform(m_whv, stddev, mu); - uniform(m_whc, 0.01*stddev, mu); uniform(m_bhv, stddev, mu); - uniform(m_bhc, 0.01*stddev, mu); uniform(m_bv, stddev, mu); - uniform(m_bc, stddev, mu); + uniform(m_ctx, stddev, mu); } void Rbm::fromJson(Json::Value rbm) @@ -74,9 +68,9 @@ Json::Value Rbm::toJson() const return rbm; } -void Rbm::gibbs_hv(arma::mat &hv_probs, arma::mat &v_probs) +void Rbm::gibbs(arma::mat &hv_probs, arma::mat &v_probs) { - for (int gibbs=0; gibbs < m_params.numGibbs; gibbs++) + for (int i=0; i < m_params.numGibbs; i++) { // Create visible reconstruction (a fantasy...) given hid if (m_params.gibbsDoSampleHidden) @@ -100,33 +94,7 @@ void Rbm::gibbs_hv(arma::mat &hv_probs, arma::mat &v_probs) } } -void Rbm::gibbs_hc(arma::mat &hc_probs, arma::mat &c_probs) -{ - for (int gibbs=0; gibbs < m_params.numGibbs; gibbs++) - { - // Create visible reconstruction (a fantasy...) given hid - if (m_params.gibbsDoSampleHidden) - { - c_probs = prob(h_to_c(sample(hc_probs))); - } - else - { - c_probs = prob(h_to_c(hc_probs)); - } - - // Create hidden representation given v - if (m_params.gibbsDoSampleVisible) - { - hc_probs = prob(c_to_h(sample(c_probs))); - } - else - { - hc_probs = prob(c_to_h(c_probs)); - } - } -} - -void Rbm::weightUpdate_hv(arma::mat const &v_states, arma::mat &dw, arma::mat &dbhv, arma::mat &dbv) +void Rbm::weightUpdate(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); @@ -147,7 +115,7 @@ void Rbm::weightUpdate_hv(arma::mat const &v_states, arma::mat &dw, arma::mat &d dbv = sum(v_states, 0); dbhv = sum(h_states, 0); - gibbs_hv(h_probs, v_probs); + gibbs(h_probs, v_probs); // Update weights (negative phase) dw -= v_probs.t() * h_probs; @@ -155,35 +123,6 @@ void Rbm::weightUpdate_hv(arma::mat const &v_states, arma::mat &dw, arma::mat &d dbhv -= sum(h_probs, 0); } -void Rbm::weightUpdate_hc(arma::mat const &c_states, arma::mat &dw, arma::mat &dbhc, arma::mat &dbc) -{ - arma::mat c_probs(c_states); - arma::mat h_states = c_to_h(c_states); - arma::mat h_probs = prob(h_states); - - // Sample hidden - if (m_params.doRaoBlackwell) - { - h_states = h_probs; - } - else - { - h_states = sample(h_probs); - } - - // Update weights (positive phase) - dw = c_states.t() * h_states; - dbc = sum(c_states, 0); - dbhc = sum(h_states, 0); - - gibbs_hc(h_probs, c_probs); - - // Update weights (negative phase) - dw -= c_probs.t() * h_probs; - dbc -= sum(c_probs, 0); - dbhc -= sum(h_probs, 0); -} - void Rbm::train(const arma::mat& batch, IListener* pListener) { Status status; @@ -193,17 +132,11 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) int batchRowIndex = 0; arma::mat grad_bias_v(arma::zeros(1, m_bv.n_cols)); - arma::mat grad_bias_c(arma::zeros(1, m_bc.n_cols)); arma::mat grad_bias_hv(arma::zeros(1, m_bhv.n_cols)); - arma::mat grad_bias_hc(arma::zeros(1, m_bhc.n_cols)); arma::mat grad_weight_hv(arma::zeros(m_whv.n_rows, m_whv.n_cols)); - arma::mat grad_weight_hc(arma::zeros(m_whc.n_rows, m_whc.n_cols)); arma::mat momentum_whv = arma::zeros(m_whv.n_rows, m_whv.n_cols); - arma::mat momentum_whc = arma::zeros(m_whc.n_rows, m_whc.n_cols); arma::mat momentum_bias_v(arma::zeros(1, m_bv.n_cols)); - arma::mat momentum_bias_c(arma::zeros(1, m_bc.n_cols)); arma::mat momentum_bias_hv(arma::zeros(1, m_bhv.n_cols)); - arma::mat momentum_bias_hc(arma::zeros(1, m_bhc.n_cols)); arma::mat penalty_weights = arma::zeros(m_whv.n_rows, m_whv.n_cols); int trainingSizeRemain = batch.n_rows; @@ -212,7 +145,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) while (trainingSizeRemain && !shouldAbort) { int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain); - arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); + arma::mat miniBatch_v = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); trainingSizeRemain -= miniBatchSizeActual; batchRowIndex += miniBatchSizeActual; int scaler = std::min(m_params.miniBatchSize, (int)batch.n_rows); @@ -224,8 +157,9 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) arma::mat v_probs(miniBatchSizeActual, m_bv.n_cols); arma::mat hid_states(miniBatchSizeActual, m_bhv.n_cols); #endif + arma::mat c_states(arma::zeros(miniBatchSizeActual, m_ctx.n_cols)); + arma::mat miniBatch(arma::join_rows(miniBatch_v, c_states)); arma::mat v_states(miniBatch); - arma::mat c_states(arma::zeros(miniBatchSizeActual, m_bc.n_cols)); // Create hidden layer base on training data if (m_params.doSampleBatch) @@ -253,7 +187,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) grad_bias_v = sum(v_states, 0); grad_bias_hv = sum(hid_states, 0); - for (int gibbs_hv=0; gibbs_hv < m_params.numGibbs; gibbs_hv++) + for (int i=0; i < m_params.numGibbs; i++) { // Create visible reconstruction (a fantasy...) given hid if (m_params.gibbsDoSampleHidden) @@ -283,14 +217,13 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) grad_bias_v -= sum(v_probs, 0); grad_bias_hv -= sum(h_probs, 0); #else - if (m_bc.n_cols == 0) + if (m_ctx.n_cols == 0) { - weightUpdate_hv(v_states, grad_weight_hv, grad_bias_hv, grad_bias_v); + weightUpdate(v_states, grad_weight_hv, grad_bias_hv, grad_bias_v); } else { - weightUpdate_hv(v_states, grad_weight_hv, grad_bias_hv, grad_bias_v); - weightUpdate_hc(c_states, grad_weight_hc, grad_bias_hc, grad_bias_c); + weightUpdate(v_states, grad_weight_hv, grad_bias_hv, grad_bias_v); } #endif @@ -300,17 +233,11 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) status.L2 = accu(m_whv % m_whv); momentum_bias_v = m_params.momentum*momentum_bias_v + grad_bias_v; momentum_bias_hv = m_params.momentum*momentum_bias_hv + grad_bias_hv; - momentum_bias_c = m_params.momentum*momentum_bias_c + grad_bias_c; - momentum_bias_hc = m_params.momentum*momentum_bias_hc + grad_bias_hc; momentum_whv = m_params.momentum*momentum_whv + grad_weight_hv - status.L2*penalty_weights; - momentum_whc = m_params.momentum*momentum_whc + grad_weight_hc; m_bv += learning_rate*momentum_bias_v; - m_bc += learning_rate*momentum_bias_c; m_bhv += learning_rate*momentum_bias_hv; - m_bhc += learning_rate*momentum_bias_hc; m_whv += learning_rate*momentum_whv; - m_whc += learning_rate*momentum_whc; progress += dProgress*miniBatchSizeActual; status.progress = (int)(progress + 0.5); @@ -337,6 +264,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) } // number of mini batches +#if FIXED_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; @@ -345,6 +273,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener) { pListener->onProgress(this, status); } +#endif } arma::mat Rbm::prob(const arma::mat &src) @@ -377,16 +306,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::c_to_h(const arma::mat &context) const -{ - return context * m_whc + arma::repmat(m_bhc, context.n_rows, 1); -} - -arma::mat Rbm::h_to_c(const arma::mat &hidden) const -{ - return hidden * m_whc.t() + arma::repmat(m_bc, hidden.n_rows, 1); -} - arma::mat Rbm::normalize(const arma::mat& src) { double mean = arma::accu(src)/src.n_elem; @@ -414,7 +333,7 @@ void Rbm::uniform(arma::mat& srcDst, double stdDev, double mu) #endif } -const arma::mat& Rbm::w() const +const arma::mat& Rbm::whv() const { return m_whv; } diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 12b2f28..5158c31 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -119,10 +119,14 @@ public: virtual ~Rbm(); void weightsInit(double stddev, double mu=0.0); + void weightsAssign(const arma::mat &w) + { + m_whv.submat(0, 0, w.n_rows-1, w.n_cols-1) = w; + } void train(arma::mat const &batch, IListener *pListener=nullptr); static arma::mat normalize(const arma::mat &hidden); - const arma::mat& w() const; + const arma::mat& whv() const; const arma::mat& bv() const; const arma::mat& bh() const; @@ -134,30 +138,35 @@ 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; - + + arma::mat toHiddenProbs(const arma::mat &visible) const + { + return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx))); + } + + arma::mat toVisibleProbs(const arma::mat &hidden) const + { + return arma::reshape(Rbm::prob(h_to_v(hidden)), 1, m_bv.n_cols-m_bhv.n_cols); + } + private: Params m_params; + arma::mat v_to_h(const arma::mat &visible) const; + arma::mat h_to_v(const arma::mat &hidden) const; arma::mat sample(arma::mat const &src); - void weightUpdate_hv(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); - void weightUpdate_hc(arma::mat const &c_states, arma::mat &dwhc, arma::mat &dbhc, arma::mat &dbc); + void weightUpdate(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_hv(arma::mat &hv_states, arma::mat &v_states); - void gibbs_hc(arma::mat &hc_states, arma::mat &c_states); + void gibbs(arma::mat &hv_states, arma::mat &v_states); protected: - arma::mat m_whv; - arma::mat m_whc; arma::mat m_bhv; - arma::mat m_bhc; arma::mat m_bv; - arma::mat m_bc; + arma::mat m_ctx; private: noise_gen_t m_noise; + arma::mat m_whv; + }; #endif /* RBM_HPP */ diff --git a/source/RbmComponent.cpp b/source/RbmComponent.cpp index 929da40..51ee66f 100644 --- a/source/RbmComponent.cpp +++ b/source/RbmComponent.cpp @@ -311,7 +311,7 @@ arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w) } void RbmComponent::redrawWeights() { - DrawWeights->getData() = getConvolutedWeight(w().col(m_currWeightIndexToDraw)); + DrawWeights->getData() = getConvolutedWeight(whv().col(m_currWeightIndexToDraw)); DrawWeights->DrawData(); }