- refactored

This commit is contained in:
2024-01-24 16:56:49 +01:00
parent 9a381c1899
commit 88f649e8f4
2 changed files with 14 additions and 14 deletions
+13 -13
View File
@@ -20,7 +20,7 @@ using namespace Matutils;
Rbm::Rbm(size_t numVisible, size_t numHidden)
: m_params()
, m_whv(numVisible, numHidden)
, m_bhv(1, numHidden)
, m_bh(1, numHidden)
, m_bv(1, numVisible)
{
assert(numVisible > 0);
@@ -30,7 +30,7 @@ Rbm::Rbm(size_t numVisible, size_t numHidden)
Rbm::Rbm(const Rbm& orig)
: m_params(orig.m_params)
, m_whv(orig.m_whv)
, m_bhv(orig.m_bhv)
, m_bh(orig.m_bh)
, m_bv(orig.m_bv)
{
}
@@ -41,7 +41,7 @@ Rbm::~Rbm()
size_t Rbm::numHidden() const
{
return m_bhv.size();
return m_bh.size();
}
size_t Rbm::numVisible() const
@@ -79,14 +79,14 @@ arma::mat Rbm::toVisibleProbs(const arma::mat& hidden) const
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_bh.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);
uniform(m_bhv, stddev, mu);
uniform(m_bh, stddev, mu);
uniform(m_bv, stddev, mu);
}
@@ -128,7 +128,7 @@ void Rbm::gibbs_hv(arma::mat &h_probs, arma::mat &v_probs) const
}
}
void Rbm::contrastiveDivergence(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 &dbh, arma::mat &dbv)
{
arma::mat v_probs(v_states);
arma::mat h_states = v_to_h(v_states);
@@ -154,7 +154,7 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::
// Update weights (positive phase)
dw = v_states.t() * h_states;
dbv = sum(v_states, 0);
dbhv = sum(h_states, 0);
dbh = sum(h_states, 0);
// Gibbs sampling with training params
for (int i=0; i < m_params.numGibbs; i++)
@@ -187,7 +187,7 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::
// Update weights (negative phase)
dw -= v_probs.t() * h_probs;
dbv -= sum(v_probs, 0);
dbhv -= sum(h_probs, 0);
dbh -= sum(h_probs, 0);
}
void Rbm::train(arma::mat const &batch, IListener* pListener)
@@ -199,11 +199,11 @@ void Rbm::train(arma::mat const &batch, IListener* pListener)
int batchRowIndex = 0;
arma::mat grad_bias_v(arma::zeros(1, m_bv.n_cols));
arma::mat grad_bias_hv(arma::zeros(1, m_bhv.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 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_bhv.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);
int trainingSizeRemain = batch.n_rows;
@@ -243,7 +243,7 @@ void Rbm::train(arma::mat const &batch, IListener* pListener)
momentum_whv = m_params.momentum*momentum_whv + grad_weight_hv - status.L2*penalty_weights;
m_bv += learning_rate*momentum_bias_v;
m_bhv += learning_rate*momentum_bias_hv;
m_bh += learning_rate*momentum_bias_hv;
m_whv += learning_rate*momentum_whv;
progress += dProgress*miniBatchSizeActual;
@@ -286,7 +286,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_bhv, 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
@@ -306,7 +306,7 @@ const arma::mat& Rbm::bv() const
const arma::mat& Rbm::bh() const
{
return m_bhv;
return m_bh;
}
+1 -1
View File
@@ -157,7 +157,7 @@ public:
protected:
Params m_params;
arma::mat m_bhv;
arma::mat m_bh;
arma::mat m_bv;
private: