diff --git a/source/AStack.cpp b/source/AStack.cpp new file mode 100644 index 0000000..573df6b --- /dev/null +++ b/source/AStack.cpp @@ -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 + +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; +} diff --git a/source/AStack.hpp b/source/AStack.hpp new file mode 100644 index 0000000..457683d --- /dev/null +++ b/source/AStack.hpp @@ -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 +#include +#include +#include +#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 */ + diff --git a/source/Layer.cpp b/source/Layer.cpp index d478b59..7e4db62 100644 --- a/source/Layer.cpp +++ b/source/Layer.cpp @@ -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) diff --git a/source/Layer.hpp b/source/Layer.hpp index d8053e5..58b6a49 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -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 */ diff --git a/source/MainComponent.cpp b/source/MainComponent.cpp index dd3358b..32f84e3 100644 --- a/source/MainComponent.cpp +++ b/source/MainComponent.cpp @@ -918,7 +918,8 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status) } RbmComponent *pComp = static_cast(pRbm); - pComp->upPass(pComp->getTraining()); + m_stack->upPass(pComp->id(), pComp->getTraining()); +// pComp->upPass(pComp->getTraining()); pComp->redrawReconstruction(); pComp->redrawWeights(); diff --git a/source/MainComponent.hpp b/source/MainComponent.hpp index 708c74f..0ffa95b 100644 --- a/source/MainComponent.hpp +++ b/source/MainComponent.hpp @@ -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(pComp); } diff --git a/source/Rbm.hpp b/source/Rbm.hpp index 3639cd5..1182186 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -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); diff --git a/source/RbmComponent.cpp b/source/RbmComponent.cpp index be203e9..62dcd61 100644 --- a/source/RbmComponent.cpp +++ b/source/RbmComponent.cpp @@ -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 (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 (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 (prev); - pComp->downPass(getReconst()); - } -} - -void RbmComponent::upDownPass(const arma::mat& vc) -{ - trainRedraw(vc); - gibbs(vc); - if (next) - { - RbmComponent *pComp = static_cast (next); - if (pComp->m_toggleEnable->getToggleState()) - { - pComp->upDownPass(DrawHidden->getData()); - } - else - { - return; - } - } - else if (prev) - { - RbmComponent *pComp = static_cast (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 (root()); - pComp->upDownPass(batch); + m_stack.upDownPass(id(), batch); } //[/MiscUserCode] diff --git a/source/RbmComponent.hpp b/source/RbmComponent.hpp index 798117d..5f283de 100644 --- a/source/RbmComponent.hpp +++ b/source/RbmComponent.hpp @@ -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 DrawVisibleTrain; ScopedPointer DrawHidden; @@ -85,8 +83,11 @@ private: ScopedPointer DrawWeights; ScopedPointer m_toggleEnable; ScopedPointer m_buttonCopyH2C; + + void onUpPass(const arma::mat& v) override; + void onDownPass(const arma::mat& h) override; - void gibbs(const arma::mat& v); + AStack &m_stack; size_t m_currWeightIndexToDraw; void onDraw(DrawComponent &obj) override; void buttonClicked(Button* buttonThatWasClicked) override;