/* * 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 "noise.h" class IRbmListener { public: virtual ~IRbmListener() { } bool onProgress() { return true; } }; class Rbm { public: struct Params { Params() : m_weightDecay(0.01) , m_learningRate(0.1) , m_momentum(0.5) , m_doRaoBlackwell(true) , m_doSampleVisible(false) , m_doSampleBatch(false) , m_numGibbs(1) { } double m_weightDecay; double m_learningRate; double m_momentum; bool m_doRaoBlackwell; bool m_doSampleVisible; bool m_doSampleBatch; size_t m_numGibbs; }; Rbm(size_t numHidden, size_t numVisible); Rbm(const Rbm& orig); virtual ~Rbm(); void train(arma::mat const &batch, size_t numEpochs, size_t sizeMiniBatch, Params const ¶ms, IRbmListener *pListener); arma::mat sample(arma::mat const &src); static arma::mat probsLogistic(arma::mat const &src); arma::mat toHidden(const arma::mat &v); arma::mat toVisible(const arma::mat &h); arma::mat uniform(size_t numRows, size_t numCols, double mu=0.0, double stdDev=1.0); void uniform(arma::mat &srcDst, double mu=0.0, double stdDev=1.0); private: noise_gen_t m_noise; arma::mat m_w; arma::mat m_bh; arma::mat m_bv; }; #endif /* RBM_HPP */