- refactored

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@811 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-15 09:19:50 +00:00
parent eb8d399636
commit 4b7c29effd
3 changed files with 66 additions and 54 deletions
+1 -4
View File
@@ -105,13 +105,10 @@ arma::mat Layer::upPass(const arma::mat& v)
arma::mat h = gibbsPass(r); arma::mat h = gibbsPass(r);
if (next) if (next)
{ {
next->upPass(h); return next->upPass(h);
} }
else
{
return h; return h;
} }
}
arma::mat Layer::upDownPass(const arma::mat& v) arma::mat Layer::upDownPass(const arma::mat& v)
{ {
+44
View File
@@ -40,6 +40,50 @@ Rbm::~Rbm()
Noise_Free(&m_noise); Noise_Free(&m_noise);
} }
size_t Rbm::numHidden() const
{
return m_bhv.size();
}
size_t Rbm::numVisible() const
{
return m_bv.size();
}
Rbm::Params& Rbm::params()
{
return m_params;
}
arma::mat Rbm::rms_error(arma::mat diffErr)
{
arma::mat diffErr_squared = diffErr % diffErr;
return arma::sum(diffErr_squared, 1) * 1.0 / diffErr_squared.n_cols;
}
double Rbm::rms_error_accu(arma::mat diffErr)
{
arma::mat diffErr_squared = diffErr % diffErr;
return arma::accu(diffErr_squared) / diffErr_squared.n_elem;
}
arma::mat Rbm::toHiddenProbs(const arma::mat& visible) const
{
return Rbm::prob(v_to_h(visible));
}
arma::mat Rbm::toVisibleProbs(const arma::mat& hidden) const
{
return Rbm::prob(h_to_v(hidden));
}
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_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);
+10 -39
View File
@@ -118,13 +118,10 @@ public:
Rbm(const Rbm& orig); Rbm(const Rbm& orig);
virtual ~Rbm(); virtual ~Rbm();
Params& params();
void weightsInit(double stddev, double mu=0.0); void weightsInit(double stddev, double mu=0.0);
void weightsAssign(const arma::mat &w, const arma::mat &bhv, const arma::mat &bv) void 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_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv;
}
void train(arma::mat const &batch, IListener *pListener=nullptr); void train(arma::mat const &batch, IListener *pListener=nullptr);
@@ -136,46 +133,20 @@ public:
Json::Value toJson() const; Json::Value toJson() const;
void fromJson(Json::Value params); void fromJson(Json::Value params);
Params& params() arma::mat toHiddenProbs(const arma::mat &visible) const;
{ arma::mat toVisibleProbs(const arma::mat &hidden) const;
return m_params;
}
static arma::mat prob(arma::mat const &src);
arma::mat toHiddenProbs(const arma::mat &visible) const size_t numHidden() const;
{
return Rbm::prob(v_to_h(visible));
}
arma::mat toVisibleProbs(const arma::mat &hidden) const size_t numVisible() const;
{
return Rbm::prob(h_to_v(hidden));
}
size_t numHidden() const
{
return m_bhv.size();
}
size_t numVisible() const
{
return m_bv.size();
}
arma::mat v_to_h(const arma::mat &visible) const; arma::mat v_to_h(const arma::mat &visible) const;
arma::mat h_to_v(const arma::mat &hidden) const; arma::mat h_to_v(const arma::mat &hidden) const;
static arma::mat prob(arma::mat const &src);
static double rms_error_accu(arma::mat diffErr) static double rms_error_accu(arma::mat diffErr);
{ static arma::mat rms_error(arma::mat diffErr);
arma::mat diffErr_squared = diffErr % diffErr;
return arma::accu(diffErr_squared)/diffErr_squared.n_elem;
}
static arma::mat rms_error(arma::mat diffErr)
{
arma::mat diffErr_squared = diffErr % diffErr;
return arma::sum(diffErr_squared, 1) * 1.0/diffErr_squared.n_cols;
}
void gibbs_hv(arma::mat &h_states, arma::mat &v_states); void gibbs_hv(arma::mat &h_states, arma::mat &v_states);
void gibbs_vh(arma::mat &v_probs, arma::mat &h_probs); void gibbs_vh(arma::mat &v_probs, arma::mat &h_probs);