- refactored
This commit is contained in:
+13
-13
@@ -20,7 +20,7 @@ using namespace Matutils;
|
|||||||
Rbm::Rbm(size_t numVisible, size_t numHidden)
|
Rbm::Rbm(size_t numVisible, size_t numHidden)
|
||||||
: m_params()
|
: m_params()
|
||||||
, m_whv(numVisible, numHidden)
|
, m_whv(numVisible, numHidden)
|
||||||
, m_bhv(1, numHidden)
|
, m_bh(1, numHidden)
|
||||||
, m_bv(1, numVisible)
|
, m_bv(1, numVisible)
|
||||||
{
|
{
|
||||||
assert(numVisible > 0);
|
assert(numVisible > 0);
|
||||||
@@ -30,7 +30,7 @@ Rbm::Rbm(size_t numVisible, size_t numHidden)
|
|||||||
Rbm::Rbm(const Rbm& orig)
|
Rbm::Rbm(const Rbm& orig)
|
||||||
: m_params(orig.m_params)
|
: m_params(orig.m_params)
|
||||||
, m_whv(orig.m_whv)
|
, m_whv(orig.m_whv)
|
||||||
, m_bhv(orig.m_bhv)
|
, m_bh(orig.m_bh)
|
||||||
, m_bv(orig.m_bv)
|
, m_bv(orig.m_bv)
|
||||||
{
|
{
|
||||||
}
|
}
|
||||||
@@ -41,7 +41,7 @@ Rbm::~Rbm()
|
|||||||
|
|
||||||
size_t Rbm::numHidden() const
|
size_t Rbm::numHidden() const
|
||||||
{
|
{
|
||||||
return m_bhv.size();
|
return m_bh.size();
|
||||||
}
|
}
|
||||||
|
|
||||||
size_t Rbm::numVisible() const
|
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)
|
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_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;
|
m_bv.submat(0, 0, bv.n_rows - 1, bv.n_cols - 1) = bv;
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::weightsInit(double stddev, double mu)
|
void Rbm::weightsInit(double stddev, double mu)
|
||||||
{
|
{
|
||||||
uniform(m_whv, stddev, mu);
|
uniform(m_whv, stddev, mu);
|
||||||
uniform(m_bhv, stddev, mu);
|
uniform(m_bh, stddev, mu);
|
||||||
uniform(m_bv, 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 v_probs(v_states);
|
||||||
arma::mat h_states = v_to_h(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)
|
// Update weights (positive phase)
|
||||||
dw = v_states.t() * h_states;
|
dw = v_states.t() * h_states;
|
||||||
dbv = sum(v_states, 0);
|
dbv = sum(v_states, 0);
|
||||||
dbhv = sum(h_states, 0);
|
dbh = sum(h_states, 0);
|
||||||
|
|
||||||
// Gibbs sampling with training params
|
// Gibbs sampling with training params
|
||||||
for (int i=0; i < m_params.numGibbs; i++)
|
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)
|
// Update weights (negative phase)
|
||||||
dw -= v_probs.t() * h_probs;
|
dw -= v_probs.t() * h_probs;
|
||||||
dbv -= sum(v_probs, 0);
|
dbv -= sum(v_probs, 0);
|
||||||
dbhv -= sum(h_probs, 0);
|
dbh -= sum(h_probs, 0);
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::train(arma::mat const &batch, IListener* pListener)
|
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;
|
int batchRowIndex = 0;
|
||||||
|
|
||||||
arma::mat grad_bias_v(arma::zeros(1, m_bv.n_cols));
|
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 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_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_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);
|
arma::mat penalty_weights = arma::zeros(m_whv.n_rows, m_whv.n_cols);
|
||||||
|
|
||||||
int trainingSizeRemain = batch.n_rows;
|
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;
|
momentum_whv = m_params.momentum*momentum_whv + grad_weight_hv - status.L2*penalty_weights;
|
||||||
|
|
||||||
m_bv += learning_rate*momentum_bias_v;
|
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;
|
m_whv += learning_rate*momentum_whv;
|
||||||
|
|
||||||
progress += dProgress*miniBatchSizeActual;
|
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
|
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
|
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
|
const arma::mat& Rbm::bh() const
|
||||||
{
|
{
|
||||||
return m_bhv;
|
return m_bh;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -157,7 +157,7 @@ public:
|
|||||||
|
|
||||||
protected:
|
protected:
|
||||||
Params m_params;
|
Params m_params;
|
||||||
arma::mat m_bhv;
|
arma::mat m_bh;
|
||||||
arma::mat m_bv;
|
arma::mat m_bv;
|
||||||
|
|
||||||
private:
|
private:
|
||||||
|
|||||||
Reference in New Issue
Block a user