- 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
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
+52 -20
View File
@@ -14,16 +14,13 @@
#include "Rbm.hpp"
#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_w(w)
, m_bv(bv)
, m_bh(bh)
, m_w(numVisible, numHidden)
, m_bv(1, numVisible)
, m_bh(1, numHidden)
{
Noise_Init(&m_noise, 0x32727155);
uniform(m_w, 0.0, m_params.weightInit);
uniform(m_bh, 0.0, m_params.weightInit);
uniform(m_bv, 0.0, m_params.weightInit);
weightsInit(0.0, m_params.weightInit);
}
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)
{
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 dst = src;
uniform(dst);
for (size_t i=0; i < src.n_rows; i++)
{
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;
}
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);
}
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);
}
arma::mat Rbm::toHiddenProbs(const arma::mat &visible)
arma::mat Rbm::toHiddenProbs(const arma::mat &visible) const
{
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));
}
void Rbm::uniform(arma::mat& srcDst, double mu, double stdDev)
{
for (size_t i=0; i < srcDst.n_rows; i++)
{
for (size_t j=0; j < srcDst.n_cols; j++)
{
srcDst(i, j) = stdDev*Noise_Uniform(&m_noise) + mu;
}
}
srcDst = stdDev*arma::randu(srcDst.n_rows, srcDst.n_cols) + 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
#include <armadillo>
#include "noise.h"
#include <jsoncpp/json/json.h>
class Rbm
{
@@ -25,7 +25,7 @@ public:
{
Params()
: weightInit(0.01)
, weightDecay(0.001)
, weightDecay(0.0)
, learningRate(0.1)
, momentum(0.5)
, 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 weightDecay;
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);
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);
arma::mat toVisibleState(const arma::mat &hidden);
arma::mat toHiddenProbs(const arma::mat &visible);
arma::mat toVisibleProbs(const arma::mat &hidden);
arma::mat& weights();
arma::mat& bias_visible();
arma::mat& bias_hidden();
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:
noise_gen_t m_noise;
const Params &m_params;
arma::mat &m_w;
arma::mat &m_bh;
arma::mat &m_bv;
arma::mat m_w;
arma::mat m_bh;
arma::mat m_bv;
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);
+18 -42
View File
@@ -7,6 +7,7 @@
#include <armadillo>
#include <jsoncpp/json/json.h>
#include "Rbm.hpp"
#include "Layer.hpp"
using namespace std;
class RbmListener : public Rbm::IListener
@@ -111,51 +112,22 @@ void saveWeight(const string &filename, size_t numVisibleX, size_t numVisibleY,
fclose(pFile);
}
void saveProject(const string &prjname)
void saveProject(const string &prjname, const Layer &layer, size_t numTraining)
{
ofstream ofs(prjname + string(".prj"));
Json::StyledWriter writer;
Json::Value project;
project["project"]["name"] = prjname;
project["project"]["num_training"] = (int)numTraining;
project["project"]["training_file"] = prjname + string(".training.dat");
Json::Value layer1;
layer1["name"] = "1";
layer1["num_hidden"] = 64;
layer1["num_visible"] = 28*28;
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 jsonLayer = layer.toJson();
Json::Value jsonLayers(Json::arrayValue);
jsonLayers.append(jsonLayer);
Json::Value layer2;
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;
project["project"]["layers"] = jsonLayers;
ofs << writer.write(project);
}
@@ -164,25 +136,29 @@ int main()
{
printf("Hallo, Welt!\n");
const string project("mnist_2");
saveProject(project);
const string project("mnist");
RbmListener statusDisplay;
Rbm::Params params;
Rbm::Params rbmParams;
arma::mat batch = loadTraining(project + string(".training.dat"));
size_t numTraining = batch.n_rows;
size_t numVisible = batch.n_cols;
size_t numHidden = 64;
Layer layer(project, 0, numVisible, numHidden);
saveProject(project, layer, numTraining);
printf("Loaded %d training samples\n", (int)numTraining);
arma::mat w = arma::zeros(numVisible, numHidden);
arma::mat bv = arma::zeros(1, numVisible);
arma::mat bh = arma::zeros(1, numHidden);
Rbm rbm(params, w, bv, bh);
Rbm rbm(rbmParams, numVisible, numHidden);
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 h = rbm.toHiddenProbs(v);
arma::mat r = rbm.toVisibleProbs(h);