Files
Rbm/source/Rbm.hpp
T
jens 0760d08516 [Rbm]
- added JSON
- refactored


git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@574 b431acfa-c32f-4a4a-93f1-934dc6c82436
2019-10-25 16:01:55 +00:00

132 lines
3.0 KiB
C++

/*
* To change this license header, choose License Headers in Project Properties.
* To change this template file, choose Tools | Templates
* and open the template in the editor.
*/
/*
* File: Rbm.hpp
* Author: jens
*
* Created on 21. Oktober 2019, 21:28
*/
#ifndef RBM_HPP
#define RBM_HPP
#include <armadillo>
#include <jsoncpp/json/json.h>
class Rbm
{
public:
struct Params
{
Params()
: weightInit(0.01)
, weightDecay(0.0)
, learningRate(0.1)
, momentum(0.5)
, doRaoBlackwell(true)
, gibbsDoSampleVisible(false)
, gibbsDoSampleHidden(true)
, doSampleBatch(false)
, numGibbs(1)
{
}
Json::Value toJson() const
{
Json::Value params;
params["weightInit"] = weightInit;
params["weightDecay"] = weightDecay;
params["learningRate"] = learningRate;
params["momentum"] = momentum;
params["doRaoBlackwell"] = (int)doRaoBlackwell;
params["gibbsDoSampleVisible"] = (int)gibbsDoSampleVisible;
params["gibbsDoSampleHidden"] = (int)gibbsDoSampleHidden;
params["doSampleBatch"] = (int)doSampleBatch;
params["numGibbs"] = (int)numGibbs;
return params;
}
void fromJson(Json::Value params)
{
weightInit = params["weightInit"].asDouble();
weightDecay = params["weightDecay"].asDouble();
learningRate = params["learningRate"].asDouble();
momentum = params["momentum"].asDouble();
doRaoBlackwell = params["doRaoBlackwell"] == 1;
gibbsDoSampleVisible = params["gibbsDoSampleVisible"] == 1;
gibbsDoSampleHidden = params["gibbsDoSampleHidden"] == 1;
doSampleBatch = params["doSampleBatch"] == 1;
numGibbs = params["numGibbs"].asUInt();
}
double weightInit;
double weightDecay;
double learningRate;
double momentum;
bool doRaoBlackwell;
bool gibbsDoSampleVisible;
bool gibbsDoSampleHidden;
bool doSampleBatch;
size_t numGibbs;
};
struct Status
{
size_t epoch;
size_t trainingSizeRemain;
double progress;
double err;
double err_total;
double L1;
double L2;
};
class IListener
{
public:
IListener() {}
virtual ~IListener() {}
virtual bool onProgress(const Status &status)
{
return true;
}
};
Rbm(const Params& params, size_t numVisible, size_t numHidden);
Rbm(const Rbm& orig);
virtual ~Rbm();
void weightsInit(double mu, double stddev);
void train(arma::mat const &batch, size_t miniBatchSize, size_t numEpochs, IListener *pListener);
arma::mat toHiddenState(const arma::mat &visible) const;
arma::mat toVisibleState(const arma::mat &hidden) const;
arma::mat toHiddenProbs(const arma::mat &visible) const;
arma::mat toVisibleProbs(const arma::mat &hidden) const;
const arma::mat& w() const;
const arma::mat& bv() const;
const arma::mat& bh() const;
Json::Value toJson() const;
void fromJson(Json::Value params);
private:
const Params &m_params;
arma::mat m_w;
arma::mat m_bh;
arma::mat m_bv;
arma::mat sample(arma::mat const &src);
static arma::mat probsLogistic(arma::mat const &src);
void uniform(arma::mat &srcDst, double mu=0.0, double stdDev=1.0);
};
#endif /* RBM_HPP */