- added JSON
- refactored


git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@574 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-25 16:01:55 +00:00
parent 836f933530
commit 0760d08516
4 changed files with 117 additions and 77 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
CONFIG ?= release CONFIG ?= release
SRCS := source/main.cpp source/Rbm.cpp source/noise.c SRCS := source/main.cpp source/Rbm.cpp source/Layer.cpp source/noise.c
LIBS := -larmadillo -ljsoncpp LIBS := -larmadillo -ljsoncpp
+52 -20
View File
@@ -14,16 +14,13 @@
#include "Rbm.hpp" #include "Rbm.hpp"
#include "noise.h" #include "noise.h"
Rbm::Rbm(const Params& params, arma::mat &w, arma::mat &bv, arma::mat &bh) Rbm::Rbm(const Params& params, size_t numVisible, size_t numHidden)
: m_params(params) : m_params(params)
, m_w(w) , m_w(numVisible, numHidden)
, m_bv(bv) , m_bv(1, numVisible)
, m_bh(bh) , m_bh(1, numHidden)
{ {
Noise_Init(&m_noise, 0x32727155); weightsInit(0.0, m_params.weightInit);
uniform(m_w, 0.0, m_params.weightInit);
uniform(m_bh, 0.0, m_params.weightInit);
uniform(m_bv, 0.0, m_params.weightInit);
} }
Rbm::Rbm(const Rbm& orig) Rbm::Rbm(const Rbm& orig)
@@ -38,6 +35,28 @@ Rbm::~Rbm()
{ {
} }
void Rbm::weightsInit(double mu, double stddev)
{
uniform(m_w, mu, stddev);
uniform(m_bh, mu, stddev);
uniform(m_bv, mu, stddev);
}
void Rbm::fromJson(Json::Value params)
{
}
Json::Value Rbm::toJson() const
{
Json::Value rbm;
rbm["numVisible"] = m_bv.n_elem;
rbm["numHidden"] = m_bh.n_elem;
rbm["params"] = m_params.toJson();
return rbm;
}
void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener) void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener)
{ {
size_t epoch; size_t epoch;
@@ -188,43 +207,56 @@ arma::mat Rbm::probsLogistic(const arma::mat &src)
arma::mat Rbm::sample(const arma::mat &src) arma::mat Rbm::sample(const arma::mat &src)
{ {
arma::mat dst = src; arma::mat dst = src;
uniform(dst);
for (size_t i=0; i < src.n_rows; i++) for (size_t i=0; i < src.n_rows; i++)
{ {
for (size_t j=0; j < src.n_cols; j++) for (size_t j=0; j < src.n_cols; j++)
{ {
dst(i, j) = src(i, j) >= Noise_Uniform(&m_noise); dst(i, j) = src(i, j) >= dst(i, j);
} }
} }
return dst; return dst;
} }
arma::mat Rbm::toHiddenState(const arma::mat &visible) arma::mat Rbm::toHiddenState(const arma::mat &visible) const
{ {
return visible * m_w + arma::repmat(m_bh, visible.n_rows, 1); return visible * m_w + arma::repmat(m_bh, visible.n_rows, 1);
} }
arma::mat Rbm::toVisibleState(const arma::mat &hidden) arma::mat Rbm::toVisibleState(const arma::mat &hidden) const
{ {
return hidden * m_w.t() + arma::repmat(m_bv, hidden.n_rows, 1); return hidden * m_w.t() + arma::repmat(m_bv, hidden.n_rows, 1);
} }
arma::mat Rbm::toHiddenProbs(const arma::mat &visible) arma::mat Rbm::toHiddenProbs(const arma::mat &visible) const
{ {
return probsLogistic(toHiddenState(visible)); return probsLogistic(toHiddenState(visible));
} }
arma::mat Rbm::toVisibleProbs(const arma::mat &hidden) arma::mat Rbm::toVisibleProbs(const arma::mat &hidden) const
{ {
return probsLogistic(toVisibleState(hidden)); return probsLogistic(toVisibleState(hidden));
} }
void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev) void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev)
{ {
for (size_t i=0; i < srcDst.n_rows; i++) srcDst = stdDev*arma::randu(srcDst.n_rows, srcDst.n_cols) + mu;
{
for (size_t j=0; j < srcDst.n_cols; j++)
{
srcDst(i, j) = stdDev*Noise_Uniform(&m_noise) + mu;
}
}
} }
const arma::mat& Rbm::w() const
{
return m_w;
}
const arma::mat& Rbm::bv() const
{
return m_bv;
}
const arma::mat& Rbm::bh() const
{
return m_bh;
}
+46 -14
View File
@@ -15,7 +15,7 @@
#define RBM_HPP #define RBM_HPP
#include <armadillo> #include <armadillo>
#include "noise.h" #include <jsoncpp/json/json.h>
class Rbm class Rbm
{ {
@@ -25,7 +25,7 @@ public:
{ {
Params() Params()
: weightInit(0.01) : weightInit(0.01)
, weightDecay(0.001) , weightDecay(0.0)
, learningRate(0.1) , learningRate(0.1)
, momentum(0.5) , momentum(0.5)
, doRaoBlackwell(true) , doRaoBlackwell(true)
@@ -36,6 +36,35 @@ public:
{ {
} }
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 weightInit;
double weightDecay; double weightDecay;
double learningRate; double learningRate;
@@ -70,26 +99,29 @@ public:
} }
}; };
Rbm(const Params& params, arma::mat &w, arma::mat &bv, arma::mat &bh); Rbm(const Params& params, size_t numVisible, size_t numHidden);
Rbm(const Rbm& orig); Rbm(const Rbm& orig);
virtual ~Rbm(); virtual ~Rbm();
void weightsInit(double mu, double stddev);
void train(arma::mat const &batch, size_t miniBatchSize, size_t numEpochs, IListener *pListener); void train(arma::mat const &batch, size_t miniBatchSize, size_t numEpochs, IListener *pListener);
arma::mat toHiddenState(const arma::mat &visible); arma::mat toHiddenState(const arma::mat &visible) const;
arma::mat toVisibleState(const arma::mat &hidden); arma::mat toVisibleState(const arma::mat &hidden) const;
arma::mat toHiddenProbs(const arma::mat &visible); arma::mat toHiddenProbs(const arma::mat &visible) const;
arma::mat toVisibleProbs(const arma::mat &hidden); arma::mat toVisibleProbs(const arma::mat &hidden) const;
arma::mat& weights(); const arma::mat& w() const;
arma::mat& bias_visible(); const arma::mat& bv() const;
arma::mat& bias_hidden(); const arma::mat& bh() const;
Json::Value toJson() const;
void fromJson(Json::Value params);
private: private:
noise_gen_t m_noise;
const Params &m_params; const Params &m_params;
arma::mat &m_w; arma::mat m_w;
arma::mat &m_bh; arma::mat m_bh;
arma::mat &m_bv; arma::mat m_bv;
arma::mat sample(arma::mat const &src); arma::mat sample(arma::mat const &src);
static arma::mat probsLogistic(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); void uniform(arma::mat &srcDst, double mu=0.0, double stdDev=1.0);
+18 -42
View File
@@ -7,6 +7,7 @@
#include <armadillo> #include <armadillo>
#include <jsoncpp/json/json.h> #include <jsoncpp/json/json.h>
#include "Rbm.hpp" #include "Rbm.hpp"
#include "Layer.hpp"
using namespace std; using namespace std;
class RbmListener : public Rbm::IListener class RbmListener : public Rbm::IListener
@@ -111,51 +112,22 @@ void saveWeight(const string &filename, size_t numVisibleX, size_t numVisibleY,
fclose(pFile); fclose(pFile);
} }
void saveProject(const string &prjname) void saveProject(const string &prjname, const Layer &layer, size_t numTraining)
{ {
ofstream ofs(prjname + string(".prj")); ofstream ofs(prjname + string(".prj"));
Json::StyledWriter writer; Json::StyledWriter writer;
Json::Value project; Json::Value project;
project["project"]["name"] = prjname; project["project"]["name"] = prjname;
project["project"]["num_training"] = (int)numTraining;
project["project"]["training_file"] = prjname + string(".training.dat");
Json::Value layer1; Json::Value jsonLayer = layer.toJson();
layer1["name"] = "1";
layer1["num_hidden"] = 64; Json::Value jsonLayers(Json::arrayValue);
layer1["num_visible"] = 28*28; jsonLayers.append(jsonLayer);
Json::Value params1;
params1["weightInit"] = 0.01;
params1["weightDecay"] = 0.001;
params1["learningRate"] = 0.1;
params1["momentum"] = 0.5;
params1["doRaoBlackwell"] = 1;
params1["gibbsDoSampleVisible"] = 0;
params1["gibbsDoSampleHidden"] = 1;
params1["doSampleBatch"] = 0;
params1["numGibbs"] = 1;
layer1["params"] = params1;
Json::Value layer2; project["project"]["layers"] = jsonLayers;
layer2["name"] = "2";
layer2["num_hidden"] = 16;
layer2["num_visible"] = 64;
Json::Value params2;
params2["weightInit"] = 0.01;
params2["weightDecay"] = 0.001;
params2["learningRate"] = 0.1;
params2["momentum"] = 0.5;
params2["doRaoBlackwell"] = 1;
params2["gibbsDoSampleVisible"] = 0;
params2["gibbsDoSampleHidden"] = 1;
params2["doSampleBatch"] = 0;
params2["numGibbs"] = 1;
layer2["params"] = params2;
Json::Value layers(Json::arrayValue);
layers.append(layer1);
layers.append(layer2);
project["project"]["layers"] = layers;
ofs << writer.write(project); ofs << writer.write(project);
} }
@@ -164,25 +136,29 @@ int main()
{ {
printf("Hallo, Welt!\n"); printf("Hallo, Welt!\n");
const string project("mnist_2"); const string project("mnist");
saveProject(project);
RbmListener statusDisplay; RbmListener statusDisplay;
Rbm::Params params; Rbm::Params rbmParams;
arma::mat batch = loadTraining(project + string(".training.dat")); arma::mat batch = loadTraining(project + string(".training.dat"));
size_t numTraining = batch.n_rows; size_t numTraining = batch.n_rows;
size_t numVisible = batch.n_cols; size_t numVisible = batch.n_cols;
size_t numHidden = 64; size_t numHidden = 64;
Layer layer(project, 0, numVisible, numHidden);
saveProject(project, layer, numTraining);
printf("Loaded %d training samples\n", (int)numTraining); printf("Loaded %d training samples\n", (int)numTraining);
arma::mat w = arma::zeros(numVisible, numHidden); arma::mat w = arma::zeros(numVisible, numHidden);
arma::mat bv = arma::zeros(1, numVisible); arma::mat bv = arma::zeros(1, numVisible);
arma::mat bh = arma::zeros(1, numHidden); arma::mat bh = arma::zeros(1, numHidden);
Rbm rbm(params, w, bv, bh); Rbm rbm(rbmParams, numVisible, numHidden);
rbm.train(batch, 1000, 100, &statusDisplay); rbm.train(batch, 1000, 100, &statusDisplay);
saveWeight(project + string(".weights.dat"), 28, 28, w, bv, bh); saveWeight(project + string(".weights.dat"), 28, 28, rbm.w(), rbm.bv(), rbm.bh());
arma::mat v = arma::randu(numTraining, numVisible); arma::mat v = arma::randu(numTraining, numVisible);
arma::mat h = rbm.toHiddenProbs(v); arma::mat h = rbm.toHiddenProbs(v);
arma::mat r = rbm.toVisibleProbs(h); arma::mat r = rbm.toVisibleProbs(h);