/* * 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 #include #include class Rbm { public: struct Params { Params() : learningRate(0.1) , weightDecay(0.0) , momentum(0.5) , doGaussianHidden(false) , doGaussianVisible(false) , doRaoBlackwell(true) , gibbsDoSampleVisible(false) , gibbsDoSampleHidden(true) , doSampleBatch(false) , numGibbs(1) , miniBatchSize(100) , numEpochs(1000) { } Json::Value toJson() const { std::cout << "Exporting Rbm::Params" << std::endl; Json::Value params; params["weightDecay"] = weightDecay; params["learningRate"] = learningRate; params["momentum"] = momentum; params["doGaussianVisible"] = (int)doGaussianVisible; params["doGaussianHidden"] = (int)doGaussianHidden; params["doRaoBlackwell"] = (int)doRaoBlackwell; params["gibbsDoSampleVisible"] = (int)gibbsDoSampleVisible; params["gibbsDoSampleHidden"] = (int)gibbsDoSampleHidden; params["doSampleBatch"] = (int)doSampleBatch; params["numGibbs"] = (int)numGibbs; params["miniBatchSize"] = (int)miniBatchSize; params["numEpochs"] = (int)numEpochs; return params; } void fromJson(Json::Value params) { std::cout << "Importing Rbm::Params" << std::endl; weightDecay = params.get("weightDecay", weightDecay).asDouble(); learningRate = params.get("learningRate", learningRate).asDouble(); momentum = params.get("momentum", momentum).asDouble(); doGaussianVisible = params.get("doGaussianVisible", doGaussianVisible) == 1; doGaussianHidden = params.get("doGaussianHidden", doGaussianHidden) == 1; doRaoBlackwell = params.get("doRaoBlackwell", doRaoBlackwell) == 1; gibbsDoSampleVisible = params.get("gibbsDoSampleVisible", gibbsDoSampleVisible) == 1; gibbsDoSampleHidden = params.get("gibbsDoSampleHidden", gibbsDoSampleHidden) == 1; doSampleBatch = params.get("doSampleBatch", doSampleBatch) == 1; numGibbs = params.get("numGibbs", numGibbs).asUInt(); miniBatchSize = params.get("miniBatchSize", miniBatchSize).asUInt(); numEpochs = params.get("numEpochs", numEpochs).asUInt(); } double weightDecay; double learningRate; double momentum; bool doGaussianVisible; bool doGaussianHidden; bool doRaoBlackwell; bool gibbsDoSampleVisible; bool gibbsDoSampleHidden; bool doSampleBatch; int numGibbs; int miniBatchSize; int numEpochs; }; struct Status { Status() : progress(0) , err(-1.0) , err_total(-1.0) , L1(-1.0) , L2(-1.0) { } int progress; double err; double err_total; double L1; double L2; }; class IListener { public: IListener() {} virtual ~IListener() {} virtual bool onProgress(Rbm *pRbm, const Status &status) { return true; } }; Rbm(size_t numVisible, size_t numHidden); 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); void train(arma::mat const &batch, IListener *pListener=nullptr); const arma::mat& whv() const; const arma::mat& bv() const; const arma::mat& bh() const; Json::Value toJson() const; void fromJson(Json::Value params); arma::mat toHiddenProbs(const arma::mat &visible) const; arma::mat toVisibleProbs(const arma::mat &hidden) const; size_t numHidden() const; size_t numVisible() const; 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); void gibbs_hv(arma::mat &h_probs, arma::mat &v_probs) const; void gibbs_vh(arma::mat &v_probs, arma::mat &h_probs) const; static double rms_error_accu(arma::mat diffErr); static arma::mat rms_error(arma::mat diffErr); protected: Params m_params; arma::mat m_bh; arma::mat m_bv; private: arma::mat m_whv; void cd_hinton_hid_binary(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); void cd_hinton_hid_linear(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); void cd_jens(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); void cd(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv); }; #endif /* RBM_HPP */