- Layer is an Rbm

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@581 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-10-26 05:59:20 +00:00
parent 0463a64a93
commit eed6956971
3 changed files with 28 additions and 39 deletions
+16 -24
View File
@@ -14,20 +14,22 @@
#include "Layer.hpp" #include "Layer.hpp"
using namespace std; using namespace std;
Layer::Layer(const string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) Layer::Layer(const string &prjname, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, const Rbm::Params &params)
: upper(nullptr) : Rbm(params, numVisibleX*numVisibleY, numHidden)
, upper(nullptr)
, lower(nullptr) , lower(nullptr)
, m_prjname(prjname) , m_prjname(prjname)
, m_id(id) , m_id(id)
, m_numVisibleX(numVisibleX) , m_numVisibleX(numVisibleX)
, m_numVisibleY(numVisibleY) , 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"); m_weightsFile = m_prjname + string(".weights.") + to_string((int)m_id) + string(".dat");
} }
Layer::Layer(const Layer& orig) Layer::Layer(const Layer& orig)
: upper(nullptr) : Rbm(orig.m_rbm_params, orig.bv().n_elem, orig.bh().n_elem)
, upper(nullptr)
, lower(nullptr) , lower(nullptr)
, m_prjname(orig.m_prjname) , m_prjname(orig.m_prjname)
, m_weightsFile(orig.m_weightsFile) , m_weightsFile(orig.m_weightsFile)
@@ -35,7 +37,6 @@ Layer::Layer(const Layer& orig)
, 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) , m_rbm_params(orig.m_rbm_params)
, m_rbm(orig.m_rbm)
{ {
} }
@@ -53,29 +54,25 @@ void Layer::saveWeights()
return; return;
} }
const arma::mat &bv = rbm().bv(); size_t numHidden = bh().n_elem;
const arma::mat &bh = rbm().bh(); size_t numVisible = bh().n_elem;
const arma::mat &w = rbm().w();
size_t numHidden = bh.n_elem;
size_t numVisible = bv.n_elem;
fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, (int)numHidden); fprintf(pFile, "%d %d %d\n", (int)m_numVisibleX, (int)m_numVisibleY, (int)numHidden);
uint32_t i, j; uint32_t i, j;
for (i=0; i < numVisible; i++) 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++) 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 (i=0; i < numVisible; i++)
{ {
for (j=0; j < numHidden; j++) for (j=0; j < numHidden; j++)
{ {
fprintf(pFile, "%3.6f ", w(i,j)); fprintf(pFile, "%3.6f ", w()(i,j));
} }
fprintf(pFile, "\n"); fprintf(pFile, "\n");
} }
@@ -84,11 +81,6 @@ void Layer::saveWeights()
} }
Rbm& Layer::rbm()
{
return m_rbm;
}
Json::Value Layer::toJson() const Json::Value Layer::toJson() const
{ {
Json::Value layer; Json::Value layer;
@@ -96,7 +88,7 @@ Json::Value Layer::toJson() const
layer["weights_file"] = m_weightsFile; layer["weights_file"] = m_weightsFile;
layer["numVisibleX"] = to_string((int)m_numVisibleX); layer["numVisibleX"] = to_string((int)m_numVisibleX);
layer["numVisibleY"] = to_string((int)m_numVisibleX); layer["numVisibleY"] = to_string((int)m_numVisibleX);
layer["rbm"] = m_rbm.toJson(); layer["rbm"] = Rbm::toJson();
return layer; return layer;
@@ -104,20 +96,20 @@ Json::Value Layer::toJson() const
arma::mat Layer::up_pass(const arma::mat &hidden) arma::mat Layer::up_pass(const arma::mat &hidden)
{ {
arma::mat reconstruction = rbm().toVisibleProbs(hidden); arma::mat reconstruction = toVisibleProbs(hidden);
if (upper) if (upper)
{ {
return upper->up_pass(reconstruction); return upper->up_pass(reconstruction);
} }
return rbm().toHiddenProbs(reconstruction); return toHiddenProbs(reconstruction);
} }
arma::mat Layer::down_pass(const arma::mat &visible) arma::mat Layer::down_pass(const arma::mat &visible)
{ {
arma::mat hidden = rbm().toHiddenProbs(visible); arma::mat hidden = toHiddenProbs(visible);
if (lower) if (lower)
{ {
return lower->down_pass(hidden); return lower->down_pass(hidden);
} }
return rbm().toVisibleProbs(hidden); return toVisibleProbs(hidden);
} }
+3 -5
View File
@@ -21,18 +21,17 @@
#include "Rbm.hpp" #include "Rbm.hpp"
#include "ILayer.hpp" #include "ILayer.hpp"
class Layer class Layer : public Rbm
{ {
public: public:
Layer *upper; Layer *upper;
Layer *lower; 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 &params);
Layer(const Layer& orig); Layer(const Layer& orig);
virtual ~Layer(); virtual ~Layer();
Json::Value toJson() const; Json::Value toJson() const;
Rbm& rbm();
void saveWeights(); void saveWeights();
arma::mat up_pass(const arma::mat& hidden); arma::mat up_pass(const arma::mat& hidden);
@@ -49,8 +48,7 @@ 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;
Rbm::Params m_rbm_params; const Rbm::Params &m_rbm_params;
Rbm m_rbm;
}; };
+9 -10
View File
@@ -93,22 +93,21 @@ int main()
printf("Loaded %d training samples\n", (int)numTraining); printf("Loaded %d training samples\n", (int)numTraining);
Layer layer(project, 0, numVisibleX, numVisibleY, numHidden); Layer layer0(project, 0, numVisibleX, numVisibleY, numHidden, rbmParams);
Layer layer1(project, 1, numVisibleX, numVisibleY, numHidden); Layer layer1(project, 1, numVisibleX, numVisibleY, numHidden, rbmParams);
Layer layer2(project, 2, numVisibleX, numVisibleY, numHidden); Layer layer2(project, 2, numVisibleX, numVisibleY, numHidden, rbmParams);
Layer layer3(project, 3, numVisibleX, numVisibleY, numHidden); Layer layer3(project, 3, numVisibleX, numVisibleY, numHidden, rbmParams);
Stack stack(project); Stack stack(project);
stack.addLayer(&layer); stack.addLayer(&layer0);
stack.addLayer(&layer1); stack.addLayer(&layer1);
stack.addLayer(&layer2); stack.addLayer(&layer2);
stack.addLayer(&layer3); stack.addLayer(&layer3);
stack.save(numTraining); stack.save(numTraining);
Rbm &rbm = layer.rbm(); layer0.train(batch, 1000, 100, &statusDisplay);
rbm.train(batch, 1000, 100, &statusDisplay); layer0.saveWeights();
layer.saveWeights();
arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY); arma::mat v = arma::randu(numTraining, numVisibleX*numVisibleY);
arma::mat h = rbm.toHiddenProbs(v); arma::mat h = layer0.toHiddenProbs(v);
arma::mat r = rbm.toVisibleProbs(h); arma::mat r = layer0.toVisibleProbs(h);
return 0; return 0;
} }