diff --git a/source/Layer.cpp b/source/Layer.cpp index 5271824..3ce69ba 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -14,20 +14,22 @@ #include "Layer.hpp" using namespace std; -Layer::Layer(const string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) -: upper(nullptr) +Layer::Layer(const string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params ¶ms) +: Rbm(params, numVisibleX*numVisibleY, numHidden) +, upper(nullptr) , lower(nullptr) , m_prjname(prjname) , m_id(id) , m_numVisibleX(numVisibleX) , m_numVisibleY(numVisibleY) -, m_rbm(m_rbm_params, numVisibleX*numVisibleY, numHidden) +, m_rbm_params(params) { m_weightsFile = m_prjname + string(".weights.") + to_string((int)m_id) + string(".dat"); } Layer::Layer(const Layer& orig) -: upper(nullptr) +: Rbm(orig.m_rbm_params, orig.bv().n_elem, orig.bh().n_elem) +, upper(nullptr) , lower(nullptr) , m_prjname(orig.m_prjname) , m_weightsFile(orig.m_weightsFile) @@ -35,7 +37,6 @@ Layer::Layer(const Layer& orig) , m_numVisibleX(orig.m_numVisibleX) , m_numVisibleY(orig.m_numVisibleY) , m_rbm_params(orig.m_rbm_params) -, m_rbm(orig.m_rbm) { } @@ -53,29 +54,25 @@ void Layer::saveWeights() return; } - const arma::mat &bv = rbm().bv(); - const arma::mat &bh = rbm().bh(); - const arma::mat &w = rbm().w(); - - size_t numHidden = bh.n_elem; - size_t numVisible = bv.n_elem; + size_t numHidden = bh().n_elem; + size_t numVisible = bh().n_elem; fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, (int)numHidden); uint32_t i, j; for (i=0; i < numVisible; i++) { - fprintf(pFile, "%3.6f\n", bv(i)); + fprintf(pFile, "%3.6f\n", bv()(i)); } for (i=0; i < numHidden; i++) { - fprintf(pFile, "%3.6f\n", bh(i)); + fprintf(pFile, "%3.6f\n", bh()(i)); } for (i=0; i < numVisible; i++) { for (j=0; j < numHidden; j++) { - fprintf(pFile, "%3.6f ", w(i,j)); + fprintf(pFile, "%3.6f ", w()(i,j)); } fprintf(pFile, "\n"); } @@ -84,11 +81,6 @@ void Layer::saveWeights() } -Rbm& Layer::rbm() -{ - return m_rbm; -} - Json::Value Layer::toJson() const { Json::Value layer; @@ -96,7 +88,7 @@ Json::Value Layer::toJson() const layer["weights_file"] = m_weightsFile; layer["numVisibleX"] = to_string((int)m_numVisibleX); layer["numVisibleY"] = to_string((int)m_numVisibleX); - layer["rbm"] = m_rbm.toJson(); + layer["rbm"] = Rbm::toJson(); return layer; @@ -104,20 +96,20 @@ Json::Value Layer::toJson() const arma::mat Layer::up_pass(const arma::mat &hidden) { - arma::mat reconstruction = rbm().toVisibleProbs(hidden); + arma::mat reconstruction = toVisibleProbs(hidden); if (upper) { return upper->up_pass(reconstruction); } - return rbm().toHiddenProbs(reconstruction); + return toHiddenProbs(reconstruction); } arma::mat Layer::down_pass(const arma::mat &visible) { - arma::mat hidden = rbm().toHiddenProbs(visible); + arma::mat hidden = toHiddenProbs(visible); if (lower) { return lower->down_pass(hidden); } - return rbm().toVisibleProbs(hidden); + return toVisibleProbs(hidden); } diff --git a/source/Layer.hpp b/source/Layer.hpp index 287bbdd..efb2adb 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -21,18 +21,17 @@ #include "Rbm.hpp" #include "ILayer.hpp" -class Layer +class Layer : public Rbm { public: Layer *upper; Layer *lower; - Layer(const std::string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden); + Layer(const std::string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params ¶ms); Layer(const Layer& orig); virtual ~Layer(); Json::Value toJson() const; - Rbm& rbm(); void saveWeights(); arma::mat up_pass(const arma::mat& hidden); @@ -49,8 +48,7 @@ private: size_t m_id; size_t m_numVisibleX; size_t m_numVisibleY; - Rbm::Params m_rbm_params; - Rbm m_rbm; + const Rbm::Params &m_rbm_params; }; diff --git a/source/main.cpp b/source/main.cpp index 973d95f..0aceea8 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -93,22 +93,21 @@ int main() printf("Loaded %d training samples\n", (int)numTraining); - Layer layer(project, 0, numVisibleX, numVisibleY, numHidden); - Layer layer1(project, 1, numVisibleX, numVisibleY, numHidden); - Layer layer2(project, 2, numVisibleX, numVisibleY, numHidden); - Layer layer3(project, 3, numVisibleX, numVisibleY, numHidden); + Layer layer0(project, 0, numVisibleX, numVisibleY, numHidden, rbmParams); + Layer layer1(project, 1, numVisibleX, numVisibleY, numHidden, rbmParams); + Layer layer2(project, 2, numVisibleX, numVisibleY, numHidden, rbmParams); + Layer layer3(project, 3, numVisibleX, numVisibleY, numHidden, rbmParams); Stack stack(project); - stack.addLayer(&layer); + stack.addLayer(&layer0); stack.addLayer(&layer1); stack.addLayer(&layer2); stack.addLayer(&layer3); stack.save(numTraining); - Rbm &rbm = layer.rbm(); - rbm.train(batch, 1000, 100, &statusDisplay); - layer.saveWeights(); + layer0.train(batch, 1000, 100, &statusDisplay); + layer0.saveWeights(); arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY); - arma::mat h = rbm.toHiddenProbs(v); - arma::mat r = rbm.toVisibleProbs(h); + arma::mat h = layer0.toHiddenProbs(v); + arma::mat r = layer0.toVisibleProbs(h); return 0; }