git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@588 b431acfa-c32f-4a4a-93f1-934dc6c82436
148 lines
3.2 KiB
C++
148 lines
3.2 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 <streambuf>
|
|
|
|
#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
|
|
{
|
|
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(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 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 */
|
|
|