- refactored
This commit is contained in:
+13
-13
@@ -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
@@ -157,7 +157,7 @@ public:
|
||||
|
||||
protected:
|
||||
Params m_params;
|
||||
arma::mat m_bhv;
|
||||
arma::mat m_bh;
|
||||
arma::mat m_bv;
|
||||
|
||||
private:
|
||||
|
||||
Reference in New Issue
Block a user