/* * 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() : 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 { std::cout << "Exporting Rbm::Params" << std::endl; 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 { Status() : epoch(0) , trainingSizeRemain(0) , progress(0) , err(-1.0) , err_total(-1.0) , L1(-1.0) , L2(-1.0) { } 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(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 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); protected: arma::mat m_w; arma::mat m_bh; arma::mat m_bv; }; #endif /* RBM_HPP */