From e2bd0ee850df281b1489a70b3d80a5649c38ce17 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Mon, 28 Oct 2019 19:20:34 +0000 Subject: [PATCH] - rbm: no training params at construction git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@591 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Rbm.cpp | 4 ++-- source/Rbm.hpp | 4 ++-- source/RbmLayer.cpp | 10 ++++------ source/RbmLayer.hpp | 3 +-- source/Stack.cpp | 5 +---- source/main.cpp | 7 +++---- 6 files changed, 13 insertions(+), 20 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index dd7f0d2..c24fb19 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -13,8 +13,8 @@ #include "Rbm.hpp" -Rbm::Rbm(const Params& params, size_t numVisible, size_t numHidden) -: m_params(params) +Rbm::Rbm(size_t numVisible, size_t numHidden) +: m_params() , m_w(numVisible, numHidden) , m_bv(1, numVisible) , m_bh(1, numHidden) diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 2523c2e..a7b3bc4 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -112,7 +112,7 @@ public: } }; - Rbm(const Params& params, size_t numVisible, size_t numHidden); + Rbm(size_t numVisible, size_t numHidden); Rbm(const Rbm& orig); virtual ~Rbm(); @@ -131,7 +131,7 @@ public: void fromJson(Json::Value params); private: - const Params &m_params; + 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); diff --git a/source/RbmLayer.cpp b/source/RbmLayer.cpp index a65aeb9..8df16de 100644 --- a/source/RbmLayer.cpp +++ b/source/RbmLayer.cpp @@ -14,22 +14,21 @@ #include "RbmLayer.hpp" using namespace std; -RbmLayer::RbmLayer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params ¶ms) -: Rbm(params, numVisibleX*numVisibleY, numHidden) +RbmLayer::RbmLayer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) +: Rbm(numVisibleX*numVisibleY, numHidden) , upper(nullptr) , lower(nullptr) , m_name(name) , m_id(id) , m_numVisibleX(numVisibleX) , m_numVisibleY(numVisibleY) -, m_rbm_params(params) { cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl; m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat"); } RbmLayer::RbmLayer(const RbmLayer& orig) -: Rbm(orig.m_rbm_params, orig.bv().n_elem, orig.bh().n_elem) +: Rbm(orig.bv().n_elem, orig.bh().n_elem) , upper(nullptr) , lower(nullptr) , m_name(orig.m_name) @@ -37,7 +36,6 @@ RbmLayer::RbmLayer(const RbmLayer& orig) , m_id(orig.m_id) , m_numVisibleX(orig.m_numVisibleX) , m_numVisibleY(orig.m_numVisibleY) -, m_rbm_params(orig.m_rbm_params) { } @@ -157,7 +155,7 @@ Json::Value RbmLayer::toJson() const layer["id"] = (int)m_id; layer["weights_file"] = m_weightsFile; layer["numVisibleX"] = (int)m_numVisibleX; - layer["numVisibleY"] = (int)m_numVisibleX; + layer["numVisibleY"] = (int)m_numVisibleY; layer["numHidden"] = (int)m_bh.n_elem; layer["rbm"] = Rbm::toJson(); diff --git a/source/RbmLayer.hpp b/source/RbmLayer.hpp index 97cd4f3..c128f48 100644 --- a/source/RbmLayer.hpp +++ b/source/RbmLayer.hpp @@ -27,7 +27,7 @@ public: RbmLayer *upper; RbmLayer *lower; - RbmLayer(const std::string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params ¶ms); + RbmLayer(const std::string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden); RbmLayer(const RbmLayer& orig); virtual ~RbmLayer(); @@ -49,7 +49,6 @@ private: size_t m_id; size_t m_numVisibleX; size_t m_numVisibleY; - const Rbm::Params &m_rbm_params; }; diff --git a/source/Stack.cpp b/source/Stack.cpp index 07ebcf9..6056f45 100644 --- a/source/Stack.cpp +++ b/source/Stack.cpp @@ -85,11 +85,8 @@ bool Stack::load() int numVisibleX = layer["numVisibleX"].asInt(); int numVisibleY = layer["numVisibleY"].asInt(); int numHidden = layer["numHidden"].asInt(); - Json::Value &rbm = layer["rbm"]["params"]; - Rbm::Params params; - params.fromJson(rbm); - addLayer(new RbmLayer(layername, i, numVisibleX, numVisibleY, numHidden, params)); + addLayer(new RbmLayer(layername, i, numVisibleX, numVisibleY, numHidden)); } return true; diff --git a/source/main.cpp b/source/main.cpp index 5ebd0b5..6ffa21d 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -82,7 +82,6 @@ int main() const string project("mnist_2"); RbmListener statusDisplay; - Rbm::Params rbmParams; Stack stack(project); arma::mat batch = loadTraining(project + string(".training.dat")); @@ -94,16 +93,16 @@ int main() printf("Loaded %d training samples\n", (int)numTraining); -#if 0 +#if 1 int i=0; - RbmLayer *lowerLayer = new RbmLayer("Layer", i, numVisibleX, numVisibleY, numHidden, rbmParams); + RbmLayer *lowerLayer = new RbmLayer("Layer", i, numVisibleX, numVisibleY, numHidden); stack.addLayer(lowerLayer); numHidden >>= 1; i++; for (i; i < 4; i++) { - RbmLayer *layer = new RbmLayer("Layer", i, lowerLayer->bh().n_elem, 1, numHidden, rbmParams); + RbmLayer *layer = new RbmLayer("Layer", i, lowerLayer->bh().n_elem, 1, numHidden); lowerLayer = layer; stack.addLayer(layer); numHidden >>= 1;