- moved method of layer interaction from Layer to Stack

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@817 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-17 08:33:30 +00:00
parent c2f5083fc3
commit a372f1a30d
9 changed files with 381 additions and 117 deletions
+239
View File
@@ -0,0 +1,239 @@
/*
* To change this license header, choose License Headers in Project Properties.
* To change this template file, choose Tools | Templates
* and open the template in the editor.
*/
/*
* File: AStack.cpp
* Author: jens
*
* Created on 16. Januar 2022, 13:55
*/
#include "AStack.hpp"
#include <cassert>
using namespace std;
AStack::AStack(const std::string &dir, StackType type, const std::string &name)
: m_name(name)
, m_type(type)
, m_pLayers(nullptr)
, m_dir(dir)
{
}
AStack::AStack(const AStack& orig)
{
}
AStack::~AStack()
{
Layer *pLayer = m_pLayers;
while(pLayer)
{
Layer *pNextLayer = pLayer->next;
delete pLayer;
pLayer = pNextLayer;
}
}
void AStack::setName(const std::string& name)
{
m_name = name;
}
size_t AStack::numLayers()
{
size_t count = 0;
Layer *pLayer = m_pLayers;
while(pLayer)
{
count++;
pLayer = pLayer->next;
}
return count;
}
void AStack::addLayer(Layer *pOtherLayer)
{
if (!m_pLayers)
{
m_pLayers = pOtherLayer;
pOtherLayer->prev = nullptr;
}
else
{
Layer *pLayer = m_pLayers;
while(pLayer->next)
{
pLayer = pLayer->next;
}
pLayer->next = pOtherLayer;
pOtherLayer->prev = pLayer;
}
}
void AStack::delLayer(Layer* pLayer)
{
assert(!"Stack::delLayer: Not implemented!");
}
Layer* AStack::getLayer(size_t layerId) const
{
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (pLayer->id() == layerId)
{
return pLayer;
}
pLayer = pLayer->next;
}
return nullptr;
}
bool AStack::load(LayerConstructor *pLayerConstructor)
{
std::cout << "Importing Project " << m_name << std::endl;
ifstream ifs(m_dir + "/" + m_name + string(".prj"));
Json::Value project;
ifs >> project;
const string &name = project["stack"]["name"].asString();
Json::Value &layers = project["stack"]["layers"];
for (int i=0; i < layers.size(); i++)
{
Json::Value &layer = layers[i];
string layername = layer["name"].asString();
int numVisibleX = layer["numVisibleX"].asInt();
int numVisibleY = layer["numVisibleY"].asInt();
int numHidden = layer["numHidden"].asInt();
int numContext = layer["numContext"].asInt();
Layer *pLayer = nullptr;
if (!pLayerConstructor)
{
pLayer = new Layer(layername, i, numVisibleX, numVisibleY, numHidden, numContext);
}
else
{
pLayer = pLayerConstructor->onConstruct(layername, i, numVisibleX, numVisibleY, numHidden, numContext);
}
assert(pLayer != nullptr);
pLayer->fromJson(layer["rbm"]);
addLayer(pLayer);
}
return true;
}
bool AStack::save()
{
std::cout << "Exporting Project " << m_name << std::endl;
ofstream ofs(m_dir + "/" + m_name + string(".prj"));
Json::Value project;
project["stack"]["name"] = m_name;
project["stack"]["type_string"] = stackTypeStrings[m_type];
project["stack"]["type"] = m_type;
Json::Value layers(Json::arrayValue);
Layer *pLayer = m_pLayers;
while(pLayer)
{
layers.append(pLayer->toJson());
pLayer = pLayer->next;
}
project["stack"]["layers"] = layers;
ofs << project;
return true;
}
void AStack::weightsInit(double stddev)
{
Layer *pLayer = m_pLayers;
while(pLayer)
{
pLayer->weightsInit(stddev);
pLayer = pLayer->next;
}
}
bool AStack::loadWeights()
{
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->weightsLoad(m_dir, m_name))
{
return false;
}
pLayer = pLayer->next;
}
return true;
}
bool AStack::saveWeights()
{
Layer *pLayer = m_pLayers;
while(pLayer)
{
if (!pLayer->weightsSave(m_dir, m_name))
{
return false;
}
pLayer = pLayer->next;
}
return true;
}
arma::mat AStack::upPass(size_t layerId, const arma::mat& v)
{
arma::mat h = arma::zeros(0,0);
arma::mat tv = v;
Layer *pLayer = getLayer(layerId);
while(pLayer)
{
if (pLayer->isEnable())
{
h = pLayer->to_h_gibbs(tv);
tv = h;
}
pLayer = pLayer->next;
}
return h;
}
arma::mat AStack::downPass(size_t layerId, const arma::mat& h)
{
arma::mat v = arma::zeros(0,0);
arma::mat th = h;
Layer *pLayer = getLayer(layerId);
while(pLayer)
{
if (pLayer->isEnable())
{
v = pLayer->to_v_gibbs(th);
th = v;
}
pLayer = pLayer->prev;
}
return v;
}
arma::mat AStack::upDownPass(size_t layerId, const arma::mat& v)
{
arma::mat h = upPass(layerId, v);
arma::mat r = downPass(numLayers()-1, h);
return r;
}
+79
View File
@@ -0,0 +1,79 @@
/*
* To change this license header, choose License Headers in Project Properties.
* To change this template file, choose Tools | Templates
* and open the template in the editor.
*/
/*
* File: AStack.hpp
* Author: jens
*
* Created on 16. Januar 2022, 13:55
*/
#ifndef ASTACK_HPP
#define ASTACK_HPP
#include <string>
#include <vector>
#include <armadillo>
#include <jsoncpp/json/json.h>
#include "Layer.hpp"
class LayerConstructor
{
public:
LayerConstructor() {}
virtual ~LayerConstructor() {}
virtual Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext)
{
return nullptr;
}
};
class AStack
{
public:
enum StackType
{
None,
Deep,
Rnn,
NUM_STACKTYPES
};
const char *stackTypeStrings[NUM_STACKTYPES] = {"None", "Deep", "Rnn"};
AStack(const std::string &dir, StackType type, const std::string &name);
AStack(const AStack& orig);
virtual ~AStack();
void setName(const std::string &name);
size_t numLayers();
void addLayer(Layer *pLayer);
void delLayer(Layer *pLayer);
Layer* getLayer(size_t layerId) const;
bool load(LayerConstructor *pLayerConstructor=nullptr);
bool save();
void weightsInit(double stddev);
bool loadWeights();
bool saveWeights();
arma::mat upPass(size_t layerId, arma::mat const &v);
arma::mat downPass(size_t layerId, arma::mat const &h);
arma::mat upDownPass(size_t layerId, arma::mat const &v);
protected:
StackType m_type;
std::string m_name;
Layer *m_pLayers;
std::string m_dir;
};
#endif /* ASTACK_HPP */
+24 -39
View File
@@ -20,6 +20,7 @@ Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibl
, prev(nullptr)
, m_name(name)
, m_id(id)
, m_isEnabled(true)
, m_numVisibleX(numVisibleX)
, m_numVisibleY(numVisibleY)
, m_numContext(numContext)
@@ -34,6 +35,7 @@ Layer::Layer(const Layer& orig)
, prev(nullptr)
, m_name(orig.m_name)
, m_id(orig.m_id)
, m_isEnabled(orig.m_isEnabled)
, m_numVisibleX(orig.m_numVisibleX)
, m_numVisibleY(orig.m_numVisibleY)
{
@@ -48,6 +50,16 @@ size_t Layer::id()
return m_id;
}
bool Layer::isEnable()
{
return m_isEnabled;
}
void Layer::setEnable(bool enable)
{
m_isEnabled = enable;
}
std::string& Layer::name()
{
return m_name;
@@ -78,53 +90,26 @@ Layer* Layer::root()
return pLayer;
}
arma::mat Layer::gibbsPass(arma::mat& vr)
arma::mat Layer::to_h_gibbs(const arma::mat& v_probs)
{
arma::mat h;
for (int i = 0; i < params().numGibbs; i++)
{
h = prob(v_to_h(vr));
vr = prob(h_to_v(h));
}
onUpPass(v_probs);
arma::mat v = v_probs;
arma::mat h = toHiddenProbs(v);
gibbs_vh(v, h);
return h;
}
arma::mat Layer::downPass(const arma::mat& h)
arma::mat Layer::to_v_gibbs(const arma::mat& h_probs)
{
arma::mat v = prob(h_to_v(h));
if (prev)
{
return prev->downPass(v);
}
onDownPass(h_probs);
arma::mat h = h_probs;
arma::mat v = toVisibleProbs(h);
gibbs_hv(h, v);
return v;
}
arma::mat Layer::upPass(const arma::mat& v)
{
arma::mat r = v;
arma::mat h = gibbsPass(r);
if (next)
{
return next->upPass(h);
}
return h;
}
arma::mat Layer::upDownPass(const arma::mat& v)
{
arma::mat r = v;
arma::mat h = gibbsPass(r);
if (next)
{
next->upDownPass(h);
}
else if (prev)
{
prev->downPass(r);
}
return r;
}
arma::mat Layer::vc_to_c(const arma::mat& vc) const
{
if (m_numContext == 0)
+11 -4
View File
@@ -33,16 +33,15 @@ public:
Json::Value toJson() const;
size_t id();
bool isEnable();
void setEnable(bool enable);
std::string& name();
int numVisibleX();
int numVisibleY();
const arma::mat& context() const;
Layer *root();
arma::mat gibbsPass(arma::mat &vr);
arma::mat upPass(arma::mat const &v);
arma::mat downPass(arma::mat const &h);
arma::mat upDownPass(arma::mat const &v);
arma::mat vc_to_v(const arma::mat &vc) const;
arma::mat vc_to_c(const arma::mat &vc) const;
@@ -53,10 +52,14 @@ public:
arma::mat trainingData(arma::mat const &batch);
void train(arma::mat const &batch, IListener *pListener=nullptr);
arma::mat to_h_gibbs(const arma::mat& v_probs);
arma::mat to_v_gibbs(const arma::mat& h_probs);
private:
std::string m_name;
size_t m_id;
bool m_isEnabled;
size_t m_numVisibleX;
size_t m_numVisibleY;
size_t m_numContext;
@@ -65,6 +68,10 @@ private:
// Compatibility
std::string filePrefix(const std::string &dir, const std::string &prjname) const;
protected:
virtual void onUpPass(const arma::mat& v) {}
virtual void onDownPass(const arma::mat& h) {}
};
#endif /* RBMLAYER_HPP */
+2 -1
View File
@@ -918,7 +918,8 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
}
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
pComp->upPass(pComp->getTraining());
m_stack->upPass(pComp->id(), pComp->getTraining());
// pComp->upPass(pComp->getTraining());
pComp->redrawReconstruction();
pComp->redrawWeights();
+1 -1
View File
@@ -74,7 +74,7 @@ public:
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, numContext);
RbmComponent *pComp = new RbmComponent(*m_stack, name, id, numVisibleX, numVisibleY, numHidden, numContext);
addAndMakeVisible(pComp);
return static_cast<Layer*>(pComp);
}
+1 -1
View File
@@ -142,7 +142,7 @@ public:
arma::mat h_to_v(const arma::mat &hidden) const;
static arma::mat prob(arma::mat const &src);
void gibbs_hv(arma::mat &h_states, arma::mat &v_states);
void gibbs_hv(arma::mat &h_probs, arma::mat &v_probs);
void gibbs_vh(arma::mat &v_probs, arma::mat &h_probs);
static double rms_error_accu(arma::mat diffErr);
+17 -65
View File
@@ -23,8 +23,9 @@
//==============================================================================
RbmComponent::RbmComponent (const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext)
RbmComponent::RbmComponent (AStack &stack, 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, numContext)
, m_stack(stack)
, m_currWeightIndexToDraw(0)
, DrawVisibleTrain(nullptr)
, DrawVisibleReconst(nullptr)
@@ -211,15 +212,15 @@ void RbmComponent::onDraw(DrawComponent &obj)
{
if (&obj == DrawHidden)
{
downPass(obj.getData());
m_stack.downPass(id(), obj.getData());
}
if (&obj == DrawVisibleTrain)
{
upDownPass(getTraining());
m_stack.upDownPass(id(), getTraining());
}
if (&obj == DrawContextTrain)
{
upDownPass(getTraining());
m_stack.upDownPass(id(), getTraining());
}
}
@@ -228,6 +229,8 @@ void RbmComponent::buttonClicked(Button* buttonThatWasClicked)
if (buttonThatWasClicked == m_toggleEnable)
{
bool state = buttonThatWasClicked->getToggleState();
setEnable(state);
if (state)
{
if (prev)
@@ -249,32 +252,14 @@ void RbmComponent::buttonClicked(Button* buttonThatWasClicked)
else if (buttonThatWasClicked == m_buttonCopyH2C)
{
DrawContextTrain->getData() = DrawHidden->getData();
upDownPass(getTraining());
m_stack.upDownPass(id(), getTraining());
}
}
void RbmComponent::redrawReconstruction()
{
RbmComponent *pComp = static_cast<RbmComponent*> (root());
pComp->upDownPass(pComp->getTraining());
}
void RbmComponent::gibbs(const arma::mat& vc)
{
arma::mat r = vc;
for (int i=0; i < params().numGibbs; i++)
{
DrawHidden->getData() = prob(v_to_h(r));
r = prob(h_to_v(DrawHidden->getData()));
}
DrawVisibleReconst->getData() = vc_to_v(r);
DrawContextReconst->getData() = vc_to_c(r);
DrawVisibleReconst->DrawData();
DrawContextReconst->DrawData();
DrawHidden->DrawData();
DrawContextReconst->DrawData();
m_stack.upDownPass(pComp->id(), pComp->getTraining());
}
arma::mat RbmComponent::getTraining() const
@@ -303,55 +288,22 @@ void RbmComponent::reconstRedraw(const arma::mat& vc)
DrawContextReconst->DrawData();
}
void RbmComponent::upPass(const arma::mat& vc)
void RbmComponent::onUpPass(const arma::mat& v)
{
DrawVisibleTrain->getData() = vc_to_v(vc);
DrawVisibleTrain->getData() = vc_to_v(v);
DrawVisibleTrain->DrawData();
DrawContextTrain->getData() = vc_to_c(vc);
DrawContextTrain->getData() = vc_to_c(v);
DrawContextTrain->DrawData();
gibbs(vc);
if (next)
{
RbmComponent *pComp = static_cast<RbmComponent*> (next);
pComp->upPass(DrawHidden->getData());
}
trainRedraw(v);
}
void RbmComponent::downPass(const arma::mat& h)
void RbmComponent::onDownPass(const arma::mat& h)
{
DrawHidden->getData() = h;
DrawHidden->DrawData();
reconstRedraw(prob(h_to_v(h)));
if (prev)
{
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
pComp->downPass(getReconst());
}
}
void RbmComponent::upDownPass(const arma::mat& vc)
{
trainRedraw(vc);
gibbs(vc);
if (next)
{
RbmComponent *pComp = static_cast<RbmComponent*> (next);
if (pComp->m_toggleEnable->getToggleState())
{
pComp->upDownPass(DrawHidden->getData());
}
else
{
return;
}
}
else if (prev)
{
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
pComp->downPass(getReconst());
}
arma::mat r = prob(h_to_v(h));
reconstRedraw(r);
}
arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w)
@@ -382,7 +334,7 @@ void RbmComponent::redrawWeights(size_t index)
void RbmComponent::setTrainingData(arma::mat const& batch)
{
RbmComponent *pComp = static_cast<RbmComponent*> (root());
pComp->upDownPass(batch);
m_stack.upDownPass(id(), batch);
}
//[/MiscUserCode]
+7 -6
View File
@@ -24,6 +24,8 @@
#include "JuceHeader.h"
#include "DrawComponent.hpp"
#include "Layer.hpp"
#include "AStack.hpp"
//[/Headers]
//==============================================================================
@@ -41,7 +43,7 @@ class RbmComponent : public Component
{
public:
//==============================================================================
RbmComponent (const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext);
RbmComponent (AStack &stack, const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext);
~RbmComponent();
//==============================================================================
@@ -63,12 +65,8 @@ public:
void redrawWeights();
void redrawWeights(size_t index);
void redrawReconstruction();
void upPass(arma::mat const &v);
void downPass(arma::mat const &v);
void upDownPass(arma::mat const &v);
arma::mat getConvolutedWeight(arma::mat const &h);
ScopedPointer<DrawComponent> DrawVisibleTrain;
ScopedPointer<DrawComponent> DrawHidden;
@@ -86,7 +84,10 @@ private:
ScopedPointer<ToggleButton> m_toggleEnable;
ScopedPointer<TextButton> m_buttonCopyH2C;
void gibbs(const arma::mat& v);
void onUpPass(const arma::mat& v) override;
void onDownPass(const arma::mat& h) override;
AStack &m_stack;
size_t m_currWeightIndexToDraw;
void onDraw(DrawComponent &obj) override;
void buttonClicked(Button* buttonThatWasClicked) override;