From 34b20f4fb3e730b7430d68043abf07cc0f18bf8f Mon Sep 17 00:00:00 2001 From: jens Date: Wed, 24 Jan 2024 19:22:57 +0100 Subject: [PATCH] - refactored - switch between cd_hinton hid/linear and cd_jens using doGaussionVisible (temporary solution) --- source/Rbm.cpp | 163 +++++++++++++++++-------------------------------- source/Rbm.hpp | 5 +- 2 files changed, 60 insertions(+), 108 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index 6296bf8..9800e3c 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -128,13 +128,51 @@ void Rbm::gibbs_hv(arma::mat &h_probs, arma::mat &v_probs) const } } -#if 1 -void Rbm::contrastiveDivergence(arma::mat const &v_data, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) +void Rbm::cd(arma::mat const &v_data, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) { - arma::mat v_probs(v_data); - arma::mat h_states = v_to_h(v_data); - arma::mat h_probs; + if (m_params.doGaussianVisible) + { + if (m_params.doGaussianHidden) + { + cd_hinton_hid_linear(v_data, dw, dbh, dbv); + } + else + { + cd_hinton_hid_binary(v_data, dw, dbh, dbv); + } + } + else + { + cd_jens(v_data, dw, dbh, dbv); + } +} + +void Rbm::cd_hinton_hid_linear(arma::mat const &v_data, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) +{ + // Start positive phase + arma::mat poshidprobs = v_to_h(v_data); + arma::mat posprods = v_data.t() * poshidprobs; + arma::mat poshidact = arma::sum(poshidprobs); + arma::mat posvisact = arma::sum(v_data); + // End of positive phase + arma::mat poshidstates = poshidprobs + arma::randn(arma::size(poshidprobs)); + + // Start negative phase + arma::mat negdata = prob(h_to_v(poshidstates)); + arma::mat neghidprobs = v_to_h(negdata); + arma::mat negprods = negdata.t() * neghidprobs; + arma::mat neghidact = arma::sum(neghidprobs); + arma::mat negvisact = arma::sum(negdata); + + // Update weight deltas + dw = posprods - negprods; + dbv = posvisact - negvisact; + dbh = poshidact - neghidact; +} + +void Rbm::cd_hinton_hid_binary(arma::mat const &v_data, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) +{ // Start positive phase arma::mat poshidprobs = prob(v_to_h(v_data)); arma::mat posprods = v_data.t() * poshidprobs; @@ -157,89 +195,7 @@ void Rbm::contrastiveDivergence(arma::mat const &v_data, arma::mat &dw, arma::ma dbh = poshidact - neghidact; } -void Rbm::train(arma::mat const &batch, IListener* pListener) -{ - Status status; - double dProgress = 100.0/(batch.n_rows*m_params.numEpochs); - double progress = 0; - int lastProgress = -100; - int batchRowIndex = 0; - - arma::mat dbv(arma::zeros(1, m_bv.n_cols)); - arma::mat dbh(arma::zeros(1, m_bh.n_cols)); - arma::mat dwhv(arma::zeros(m_whv.n_rows, m_whv.n_cols)); - arma::mat inc_whv = arma::zeros(m_whv.n_rows, m_whv.n_cols); - arma::mat inc_bv(arma::zeros(1, m_bv.n_cols)); - arma::mat inc_bh(arma::zeros(1, m_bh.n_cols)); - - int trainingSizeRemain = batch.n_rows; - - bool shouldAbort = false; - while (trainingSizeRemain && !shouldAbort) - { - int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain); - arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); - trainingSizeRemain -= miniBatchSizeActual; - batchRowIndex += miniBatchSizeActual; - int numcases = std::min(m_params.miniBatchSize, (int)batch.n_rows); - - arma::mat v_states(miniBatch); - - // Create hidden layer base on training data - if (m_params.doSampleBatch) - { - // When the hidden units are being driven by data, always use stochastic binary states - v_states = sample(miniBatch); - } - - for (int epoch=0; epoch < m_params.numEpochs; epoch++) - { - // Contrastive divergence learning: calculate gradients - contrastiveDivergence(v_states, dwhv, dbh, dbv); - - // Adjust weight and biases - inc_bv = m_params.momentum*inc_bv + m_params.learningRate/numcases*dbv; - inc_bh = m_params.momentum*inc_bh + m_params.learningRate/numcases*dbh; - inc_whv = m_params.momentum*inc_whv + m_params.learningRate*(dwhv/numcases - m_params.weightDecay*m_whv); - - m_bv += inc_bv; - m_bh += inc_bh; - m_whv += inc_whv; - - progress += dProgress*miniBatchSizeActual; - status.progress = (int)(progress + 0.5); - - // Update status - if (status.progress != lastProgress) - { - lastProgress = status.progress; - - // Calculate error - status.err = rms_error_accu(miniBatch - prob(h_to_v(prob(v_to_h(v_states))))); - if (pListener) - { - if(!pListener->onProgress(this, status)) - { - shouldAbort = true; - break; - } - } - } - - } // Number of epochs - - } // number of mini batches - - // Update final status - status.err_total = rms_error_accu(batch - prob(h_to_v(prob(v_to_h(batch))))); - - if (pListener) - { - pListener->onProgress(this, status); - } -} -#else -void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) +void Rbm::cd_jens(arma::mat const &v_states, arma::mat &dw, arma::mat &dbh, arma::mat &dbv) { arma::mat v_probs(v_states); arma::mat h_states = v_to_h(v_states); @@ -300,6 +256,7 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma:: dbv -= sum(v_probs, 0); dbh -= sum(h_probs, 0); } + void Rbm::train(arma::mat const &batch, IListener* pListener) { Status status; @@ -311,10 +268,9 @@ void Rbm::train(arma::mat const &batch, IListener* pListener) arma::mat dbv(arma::zeros(1, m_bv.n_cols)); arma::mat dbh(arma::zeros(1, m_bh.n_cols)); arma::mat dwhv(arma::zeros(m_whv.n_rows, m_whv.n_cols)); - arma::mat momentum_whv = arma::zeros(m_whv.n_rows, m_whv.n_cols); - arma::mat momentum_bv(arma::zeros(1, m_bv.n_cols)); - arma::mat momentum_bh(arma::zeros(1, m_bh.n_cols)); - arma::mat penalty_whv = arma::zeros(m_whv.n_rows, m_whv.n_cols); + arma::mat inc_whv = arma::zeros(m_whv.n_rows, m_whv.n_cols); + arma::mat inc_bv(arma::zeros(1, m_bv.n_cols)); + arma::mat inc_bh(arma::zeros(1, m_bh.n_cols)); int trainingSizeRemain = batch.n_rows; @@ -325,9 +281,7 @@ void Rbm::train(arma::mat const &batch, IListener* pListener) arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); trainingSizeRemain -= miniBatchSizeActual; batchRowIndex += miniBatchSizeActual; - int scaler = std::min(m_params.miniBatchSize, (int)batch.n_rows); - double learning_rate = m_params.learningRate/scaler; - double weight_decay = m_params.weightDecay/scaler; + int numcases = std::min(m_params.miniBatchSize, (int)batch.n_rows); arma::mat v_states(miniBatch); @@ -341,20 +295,16 @@ void Rbm::train(arma::mat const &batch, IListener* pListener) for (int epoch=0; epoch < m_params.numEpochs; epoch++) { // Contrastive divergence learning: calculate gradients - contrastiveDivergence(v_states, dwhv, dbh, dbv); + cd(v_states, dwhv, dbh, dbv); // Adjust weight and biases - penalty_whv = weight_decay*arma::sign(m_whv); + inc_bv = m_params.momentum*inc_bv + m_params.learningRate/numcases*dbv; + inc_bh = m_params.momentum*inc_bh + m_params.learningRate/numcases*dbh; + inc_whv = m_params.momentum*inc_whv + m_params.learningRate*(dwhv/numcases - m_params.weightDecay*m_whv); - status.L1 = accu(abs(m_whv)); - status.L2 = accu(m_whv % m_whv); - momentum_bv = m_params.momentum*momentum_bv + dbv; - momentum_bh = m_params.momentum*momentum_bh + dbh; - momentum_whv = m_params.momentum*momentum_whv + dwhv - status.L2*penalty_whv; - - m_bv += learning_rate*momentum_bv; - m_bh += learning_rate*momentum_bh; - m_whv += learning_rate*momentum_whv; + m_bv += inc_bv; + m_bh += inc_bh; + m_whv += inc_whv; progress += dProgress*miniBatchSizeActual; status.progress = (int)(progress + 0.5); @@ -388,7 +338,6 @@ void Rbm::train(arma::mat const &batch, IListener* pListener) pListener->onProgress(this, status); } } -#endif arma::mat Rbm::prob(const arma::mat &src) { diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 59eb069..c63ced8 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -162,7 +162,10 @@ protected: private: arma::mat m_whv; - void contrastiveDivergence(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); + void cd_hinton_hid_binary(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); + void cd_hinton_hid_linear(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); + void cd_jens(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); + void cd(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); };