- pass numContext

- fixed crash when numContext == 0

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@761 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-09 08:01:25 +00:00
parent 0ee0a732e9
commit c2ffecd66a
10 changed files with 46 additions and 21 deletions
+10 -7
View File
@@ -14,8 +14,8 @@
#include "Layer.hpp" #include "Layer.hpp"
using namespace std; using namespace std;
Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext)
: Rbm(numVisibleX*numVisibleY, numHidden, numHidden) : Rbm(numVisibleX*numVisibleY, numHidden, numContext)
, next(nullptr) , next(nullptr)
, prev(nullptr) , prev(nullptr)
, m_name(name) , m_name(name)
@@ -73,23 +73,25 @@ bool Layer::loadWeights(const string &prjname)
int i, j; int i, j;
float v; float v;
arma::mat _bv(1, numVisible);
for (i=0; i < numVisible; i++) for (i=0; i < numVisible; i++)
{ {
result = fscanf(pFile, "%f", &v); result = fscanf(pFile, "%f", &v);
if (result > 0) if (result > 0)
{ {
m_bv(i) = v; _bv(i) = v;
} }
} }
arma::mat _bhv(1, numHidden);
for (i=0; i < numHidden; i++) for (i=0; i < numHidden; i++)
{ {
result = fscanf(pFile, "%f", &v); result = fscanf(pFile, "%f", &v);
if (result > 0) if (result > 0)
{ {
m_bhv(i) = v; _bhv(i) = v;
} }
} }
arma::mat weights(numVisible, numHidden); arma::mat _whv(numVisible, numHidden);
for (i=0; i < numVisible; i++) for (i=0; i < numVisible; i++)
{ {
for (j=0; j < numHidden; j++) for (j=0; j < numHidden; j++)
@@ -98,11 +100,11 @@ bool Layer::loadWeights(const string &prjname)
result = fscanf(pFile, "%f", &v); result = fscanf(pFile, "%f", &v);
if (result > 0) if (result > 0)
{ {
weights(i, j) = v; _whv(i, j) = v;
} }
} }
} }
weightsAssign(weights); weightsAssign(_whv, _bhv, _bv);
fclose(pFile); fclose(pFile);
return true; return true;
@@ -162,6 +164,7 @@ Json::Value Layer::toJson() const
layer["numVisibleX"] = (int)m_numVisibleX; layer["numVisibleX"] = (int)m_numVisibleX;
layer["numVisibleY"] = (int)m_numVisibleY; layer["numVisibleY"] = (int)m_numVisibleY;
layer["numHidden"] = (int)whv().n_cols; layer["numHidden"] = (int)whv().n_cols;
layer["numContext"] = numContext();
layer["rbm"] = Rbm::toJson(); layer["rbm"] = Rbm::toJson();
return layer; return layer;
+1 -1
View File
@@ -26,7 +26,7 @@ public:
Layer *next; Layer *next;
Layer *prev; Layer *prev;
Layer(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden); Layer(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext=0);
Layer(const Layer& orig); Layer(const Layer& orig);
virtual ~Layer(); virtual ~Layer();
+3 -1
View File
@@ -551,6 +551,8 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
int numVisX = numVisibleLabel->getText().getIntValue(); int numVisX = numVisibleLabel->getText().getIntValue();
int numVisY = numVisibleYLabel->getText().getIntValue(); int numVisY = numVisibleYLabel->getText().getIntValue();
int numHid = numHiddenLabel->getText().getIntValue(); int numHid = numHiddenLabel->getText().getIntValue();
int numCtx = 0;
if (!m_stack) if (!m_stack)
{ {
m_stack = new Stack(m_file.getParentDirectory().getFullPathName().toStdString(), std::string(projectNameLabel->getText().getCharPointer())); m_stack = new Stack(m_file.getParentDirectory().getFullPathName().toStdString(), std::string(projectNameLabel->getText().getCharPointer()));
@@ -561,7 +563,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
numVisX = pPrev->bh().n_elem; numVisX = pPrev->bh().n_elem;
numVisY = 1; numVisY = 1;
} }
m_stack->addLayer(onConstruct("Layer", next_index, numVisX, numVisY, numHid)); m_stack->addLayer(onConstruct("Layer", next_index, numVisX, numVisY, numHid, numCtx));
m_rbmSelect->addItem(String(next_index), next_index+1); m_rbmSelect->addItem(String(next_index), next_index+1);
m_rbmSelect->setSelectedId(next_index+1, sendNotification); m_rbmSelect->setSelectedId(next_index+1, sendNotification);
//[/UserButtonCode_createButton] //[/UserButtonCode_createButton]
+2 -2
View File
@@ -72,9 +72,9 @@ public:
void mouseDoubleClick (const MouseEvent& e); void mouseDoubleClick (const MouseEvent& e);
void mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel); void mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel);
Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext)
{ {
RbmComponent *pComp = new RbmComponent(name, id, numVisibleX, numVisibleY, numHidden); RbmComponent *pComp = new RbmComponent(name, id, numVisibleX, numVisibleY, numHidden, numContext);
addAndMakeVisible(pComp); addAndMakeVisible(pComp);
return static_cast<Layer*>(pComp); return static_cast<Layer*>(pComp);
} }
+2 -1
View File
@@ -21,13 +21,14 @@ Rbm::Rbm(size_t numVisible, size_t numHidden, size_t numContext)
, m_whv(numVisible+numContext, numHidden) , m_whv(numVisible+numContext, numHidden)
, m_bhv(1, numHidden) , m_bhv(1, numHidden)
, m_bv(1, numVisible+numContext) , m_bv(1, numVisible+numContext)
, m_ctx(1, numContext) , m_ctx()
{ {
assert(numVisible > 0); assert(numVisible > 0);
assert(numHidden > 0); assert(numHidden > 0);
if (numContext) if (numContext)
{ {
assert(numContext == numHidden); assert(numContext == numHidden);
m_ctx.resize(1, numContext);
} }
Noise_Init(&m_noise, 0x32727155); Noise_Init(&m_noise, 0x32727155);
} }
+21 -3
View File
@@ -119,10 +119,13 @@ public:
virtual ~Rbm(); virtual ~Rbm();
void weightsInit(double stddev, double mu=0.0); void weightsInit(double stddev, double mu=0.0);
void weightsAssign(const arma::mat &w) void weightsAssign(const arma::mat &w, const arma::mat &bhv, const arma::mat &bv)
{ {
m_whv.submat(0, 0, w.n_rows-1, w.n_cols-1) = w; m_whv.submat(0, 0, w.n_rows-1, w.n_cols-1) = w;
m_bhv.submat(0, 0, bhv.n_rows-1, bhv.n_cols-1) = bhv;
m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv;
} }
void train(arma::mat const &batch, IListener *pListener=nullptr); void train(arma::mat const &batch, IListener *pListener=nullptr);
static arma::mat normalize(const arma::mat &hidden); static arma::mat normalize(const arma::mat &hidden);
@@ -141,12 +144,27 @@ public:
arma::mat toHiddenProbs(const arma::mat &visible) const arma::mat toHiddenProbs(const arma::mat &visible) const
{ {
return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx))); return Rbm::prob(v_to_h(arma::join_rows(visible, m_ctx)));
} }
arma::mat toVisibleProbs(const arma::mat &hidden) const arma::mat toVisibleProbs(const arma::mat &hidden) const
{ {
return arma::reshape(Rbm::prob(h_to_v(hidden)), 1, m_bv.n_cols-m_bhv.n_cols); return arma::reshape(Rbm::prob(h_to_v(hidden)), 1, numVisible() - numContext());
}
size_t numContext() const
{
return m_ctx.size();
}
size_t numHidden() const
{
return m_bhv.size();
}
size_t numVisible() const
{
return m_bv.size();
} }
private: private:
+2 -2
View File
@@ -23,8 +23,8 @@
//============================================================================== //==============================================================================
RbmComponent::RbmComponent (const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) RbmComponent::RbmComponent (const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext)
: Layer(name, id, numVisibleX, numVisibleY, numHidden) : Layer(name, id, numVisibleX, numVisibleY, numHidden, numContext)
, m_currWeightIndexToDraw(0) , m_currWeightIndexToDraw(0)
, DrawTraining(nullptr) , DrawTraining(nullptr)
, DrawReconstruction(nullptr) , DrawReconstruction(nullptr)
+1 -1
View File
@@ -41,7 +41,7 @@ class RbmComponent : public Component
{ {
public: public:
//============================================================================== //==============================================================================
RbmComponent (const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden); RbmComponent (const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext);
~RbmComponent(); ~RbmComponent();
//============================================================================== //==============================================================================
+3 -2
View File
@@ -116,15 +116,16 @@ bool Stack::load(LayerConstructor *pLayerConstructor)
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();
int numContext = layer["numContext"].asInt();
Layer *pLayer = nullptr; Layer *pLayer = nullptr;
if (!pLayerConstructor) if (!pLayerConstructor)
{ {
pLayer = new Layer(layername, i, numVisibleX, numVisibleY, numHidden); pLayer = new Layer(layername, i, numVisibleX, numVisibleY, numHidden, numContext);
} }
else else
{ {
pLayer = pLayerConstructor->onConstruct(layername, i, numVisibleX, numVisibleY, numHidden); pLayer = pLayerConstructor->onConstruct(layername, i, numVisibleX, numVisibleY, numHidden, numContext);
} }
assert(pLayer != nullptr); assert(pLayer != nullptr);
+1 -1
View File
@@ -26,7 +26,7 @@ public:
LayerConstructor() {} LayerConstructor() {}
virtual ~LayerConstructor() {} virtual ~LayerConstructor() {}
virtual Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden) virtual Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext)
{ {
return nullptr; return nullptr;
} }