- refactored
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@811 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+2
-5
@@ -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)
|
||||
|
||||
@@ -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
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user