- rbm: no training params at construction

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@591 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-28 19:20:34 +00:00
parent 711ad3057e
commit e2bd0ee850
6 changed files with 13 additions and 20 deletions
+2 -2
View File
@@ -13,8 +13,8 @@
#include "Rbm.hpp" #include "Rbm.hpp"
Rbm::Rbm(const Params& params, size_t numVisible, size_t numHidden) Rbm::Rbm(size_t numVisible, size_t numHidden)
: m_params(params) : m_params()
, m_w(numVisible, numHidden) , m_w(numVisible, numHidden)
, m_bv(1, numVisible) , m_bv(1, numVisible)
, m_bh(1, numHidden) , m_bh(1, numHidden)
+2 -2
View File
@@ -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); Rbm(const Rbm& orig);
virtual ~Rbm(); virtual ~Rbm();
@@ -131,7 +131,7 @@ public:
void fromJson(Json::Value params); void fromJson(Json::Value params);
private: private:
const Params &m_params; const Params m_params;
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);
+4 -6
View File
@@ -14,22 +14,21 @@
#include "RbmLayer.hpp" #include "RbmLayer.hpp"
using namespace std; using namespace std;
RbmLayer::RbmLayer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params &params) RbmLayer::RbmLayer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden)
: Rbm(params, numVisibleX*numVisibleY, numHidden) : Rbm(numVisibleX*numVisibleY, numHidden)
, upper(nullptr) , upper(nullptr)
, lower(nullptr) , lower(nullptr)
, m_name(name) , m_name(name)
, m_id(id) , m_id(id)
, m_numVisibleX(numVisibleX) , m_numVisibleX(numVisibleX)
, m_numVisibleY(numVisibleY) , m_numVisibleY(numVisibleY)
, m_rbm_params(params)
{ {
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl; cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat"); m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
} }
RbmLayer::RbmLayer(const RbmLayer& orig) 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) , upper(nullptr)
, lower(nullptr) , lower(nullptr)
, m_name(orig.m_name) , m_name(orig.m_name)
@@ -37,7 +36,6 @@ RbmLayer::RbmLayer(const RbmLayer& orig)
, m_id(orig.m_id) , m_id(orig.m_id)
, m_numVisibleX(orig.m_numVisibleX) , m_numVisibleX(orig.m_numVisibleX)
, m_numVisibleY(orig.m_numVisibleY) , 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["id"] = (int)m_id;
layer["weights_file"] = m_weightsFile; layer["weights_file"] = m_weightsFile;
layer["numVisibleX"] = (int)m_numVisibleX; layer["numVisibleX"] = (int)m_numVisibleX;
layer["numVisibleY"] = (int)m_numVisibleX; layer["numVisibleY"] = (int)m_numVisibleY;
layer["numHidden"] = (int)m_bh.n_elem; layer["numHidden"] = (int)m_bh.n_elem;
layer["rbm"] = Rbm::toJson(); layer["rbm"] = Rbm::toJson();
+1 -2
View File
@@ -27,7 +27,7 @@ public:
RbmLayer *upper; RbmLayer *upper;
RbmLayer *lower; RbmLayer *lower;
RbmLayer(const std::string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params &params); RbmLayer(const std::string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden);
RbmLayer(const RbmLayer& orig); RbmLayer(const RbmLayer& orig);
virtual ~RbmLayer(); virtual ~RbmLayer();
@@ -49,7 +49,6 @@ private:
size_t m_id; size_t m_id;
size_t m_numVisibleX; size_t m_numVisibleX;
size_t m_numVisibleY; size_t m_numVisibleY;
const Rbm::Params &m_rbm_params;
}; };
+1 -4
View File
@@ -85,11 +85,8 @@ bool Stack::load()
int numVisibleX = layer["numVisibleX"].asInt(); int numVisibleX = layer["numVisibleX"].asInt();
int numVisibleY = layer["numVisibleY"].asInt(); int numVisibleY = layer["numVisibleY"].asInt();
int numHidden = layer["numHidden"].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; return true;
+3 -4
View File
@@ -82,7 +82,6 @@ int main()
const string project("mnist_2"); const string project("mnist_2");
RbmListener statusDisplay; RbmListener statusDisplay;
Rbm::Params rbmParams;
Stack stack(project); Stack stack(project);
arma::mat batch = loadTraining(project + string(".training.dat")); arma::mat batch = loadTraining(project + string(".training.dat"));
@@ -94,16 +93,16 @@ int main()
printf("Loaded %d training samples\n", (int)numTraining); printf("Loaded %d training samples\n", (int)numTraining);
#if 0 #if 1
int i=0; 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); stack.addLayer(lowerLayer);
numHidden >>= 1; numHidden >>= 1;
i++; i++;
for (i; i < 4; 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; lowerLayer = layer;
stack.addLayer(layer); stack.addLayer(layer);
numHidden >>= 1; numHidden >>= 1;