- 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
+2 -5
View File
@@ -105,12 +105,9 @@ arma::mat Layer::upPass(const arma::mat& v)
arma::mat h = gibbsPass(r);
if (next)
{
next->upPass(h);
}
else
{
return h;
return next->upPass(h);
}
return h;
}
arma::mat Layer::upDownPass(const arma::mat& v)
+44
View File
@@ -40,6 +40,50 @@ Rbm::~Rbm()
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)
{
uniform(m_whv, stddev, mu);
+20 -49
View File
@@ -23,7 +23,7 @@
class Rbm
{
public:
struct Params
{
Params()
@@ -72,7 +72,7 @@ public:
miniBatchSize = params.get("miniBatchSize", miniBatchSize).asUInt();
numEpochs = params.get("numEpochs", numEpochs).asUInt();
}
double weightDecay;
double learningRate;
double momentum;
@@ -110,7 +110,7 @@ public:
virtual ~IListener() {}
virtual bool onProgress(Rbm *pRbm, const Status &status)
{
return true;
return true;
}
};
@@ -118,73 +118,44 @@ public:
Rbm(const Rbm& orig);
virtual ~Rbm();
Params& params();
void weightsInit(double stddev, double mu=0.0);
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 weightsAssign(const arma::mat &w, const arma::mat &bhv, const arma::mat &bv);
void train(arma::mat const &batch, IListener *pListener=nullptr);
static arma::mat normalize(const arma::mat &hidden);
const arma::mat& whv() const;
const arma::mat& bv() const;
const arma::mat& bh() const;
Json::Value toJson() const;
void fromJson(Json::Value params);
Params& params()
{
return m_params;
}
static arma::mat prob(arma::mat const &src);
arma::mat toHiddenProbs(const arma::mat &visible) const
{
return Rbm::prob(v_to_h(visible));
}
arma::mat toVisibleProbs(const arma::mat &hidden) const
{
return Rbm::prob(h_to_v(hidden));
}
size_t numHidden() const
{
return m_bhv.size();
}
arma::mat toHiddenProbs(const arma::mat &visible) const;
arma::mat toVisibleProbs(const arma::mat &hidden) const;
size_t numHidden() const;
size_t numVisible() const;
size_t numVisible() const
{
return m_bv.size();
}
arma::mat v_to_h(const arma::mat &visible) 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)
{
arma::mat diffErr_squared = diffErr % diffErr;
return arma::accu(diffErr_squared)/diffErr_squared.n_elem;
}
static double rms_error_accu(arma::mat diffErr);
static arma::mat rms_error(arma::mat diffErr);
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_vh(arma::mat &v_probs, arma::mat &h_probs);
private:
Params m_params;
arma::mat sample(arma::mat const &src);
void contrastiveDivergence(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv);
void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5);
protected:
arma::mat m_bhv;
arma::mat m_bv;