- 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:
@@ -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;
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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 */
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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]
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user