From a5ed0be991f83ed2339c0a64a61378d3ed11b882 Mon Sep 17 00:00:00 2001 From: jens Date: Wed, 24 Jan 2024 18:41:42 +0100 Subject: [PATCH] - added original implementation for binary RBM --- source/Rbm.cpp | 145 ++++++++++++++++++++++++++++++++++++++------ source/matutils.hpp | 6 +- 2 files changed, 131 insertions(+), 20 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index ff039c5..6296bf8 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -86,8 +86,8 @@ void Rbm::weightsAssign(const arma::mat& w, const arma::mat& bh, const arma::mat void Rbm::weightsInit(double stddev, double mu) { uniform(m_whv, stddev, mu); - uniform(m_bh, stddev, mu); - uniform(m_bv, stddev, mu); + uniform(m_bh, 0, mu); + uniform(m_bv, 0, mu); } void Rbm::fromJson(Json::Value rbm) @@ -128,6 +128,117 @@ 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) +{ + arma::mat v_probs(v_data); + arma::mat h_states = v_to_h(v_data); + arma::mat h_probs; + + // Start positive phase + arma::mat poshidprobs = prob(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 = sample(poshidprobs); + + // Start negative phase + arma::mat negdata = prob(h_to_v(poshidstates)); + arma::mat neghidprobs = prob(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::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) { arma::mat v_probs(v_states); @@ -189,7 +300,6 @@ 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; @@ -198,13 +308,13 @@ void Rbm::train(arma::mat const &batch, IListener* pListener) int lastProgress = -100; int batchRowIndex = 0; - arma::mat grad_bias_v(arma::zeros(1, m_bv.n_cols)); - arma::mat grad_bias_hv(arma::zeros(1, m_bh.n_cols)); - arma::mat grad_weight_hv(arma::zeros(m_whv.n_rows, m_whv.n_cols)); + 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_bias_v(arma::zeros(1, m_bv.n_cols)); - arma::mat momentum_bias_hv(arma::zeros(1, m_bh.n_cols)); - arma::mat penalty_weights = 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); int trainingSizeRemain = batch.n_rows; @@ -231,19 +341,19 @@ 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, grad_weight_hv, grad_bias_hv, grad_bias_v); + contrastiveDivergence(v_states, dwhv, dbh, dbv); // Adjust weight and biases - penalty_weights = weight_decay*arma::sign(m_whv); + penalty_whv = weight_decay*arma::sign(m_whv); status.L1 = accu(abs(m_whv)); 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_whv = m_params.momentum*momentum_whv + grad_weight_hv - status.L2*penalty_weights; + 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_bias_v; - m_bh += learning_rate*momentum_bias_hv; + m_bv += learning_rate*momentum_bv; + m_bh += learning_rate*momentum_bh; m_whv += learning_rate*momentum_whv; progress += dProgress*miniBatchSizeActual; @@ -278,6 +388,7 @@ void Rbm::train(arma::mat const &batch, IListener* pListener) pListener->onProgress(this, status); } } +#endif arma::mat Rbm::prob(const arma::mat &src) { @@ -286,7 +397,7 @@ arma::mat Rbm::prob(const arma::mat &src) arma::mat Rbm::v_to_h(const arma::mat &visible) const { - return visible * m_whv + arma::repmat(m_bh, visible.n_rows, 1); + return visible * m_whv + arma::repmat(m_bh, visible.n_rows, 1); } arma::mat Rbm::h_to_v(const arma::mat &hidden) const diff --git a/source/matutils.hpp b/source/matutils.hpp index 02e9f27..fcac863 100644 --- a/source/matutils.hpp +++ b/source/matutils.hpp @@ -37,8 +37,8 @@ namespace Matutils inline arma::mat sample(const arma::mat &src) { - arma::mat dst = src; - uniform(dst); + arma::mat rand = src; + uniform(rand); #if 0 for (size_t i=0; i < src.n_rows; i++) @@ -50,7 +50,7 @@ namespace Matutils } return dst; #else - arma::umat res = (dst < src); + arma::umat res = (src > rand); return arma::conv_to::from(res); #endif