- fixed gui wiring
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@764 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
@@ -533,7 +533,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
||||
{
|
||||
//[UserButtonCode_addButton] -- add your button handler code here..
|
||||
RbmComponent *pComp = static_cast<RbmComponent*>(m_stack->getLayer(0));
|
||||
addTraining(pComp->DrawTraining->getData());
|
||||
addTraining(pComp->DrawVisibleTrain->getData());
|
||||
//[/UserButtonCode_addButton]
|
||||
}
|
||||
else if (buttonThatWasClicked == ShakeButton)
|
||||
@@ -915,7 +915,7 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
|
||||
}
|
||||
|
||||
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
|
||||
pComp->upPass(pComp->DrawTraining->getData());
|
||||
pComp->upPass(pComp->getTraining());
|
||||
pComp->redrawReconstruction();
|
||||
pComp->redrawWeights();
|
||||
|
||||
|
||||
@@ -313,6 +313,21 @@ arma::mat Rbm::sample(const arma::mat &src)
|
||||
return dst;
|
||||
}
|
||||
|
||||
arma::mat Rbm::vc_to_v(const arma::mat &vc) const
|
||||
{
|
||||
return arma::reshape(vc, 1, numVisible() - numContext());
|
||||
}
|
||||
|
||||
arma::mat Rbm::vc_to_c(const arma::mat &vc) const
|
||||
{
|
||||
return vc.submat(0, numVisible() - numContext(), 0, numVisible() - 1);
|
||||
}
|
||||
|
||||
arma::mat Rbm::to_vc(const arma::mat &v, const arma::mat &c) const
|
||||
{
|
||||
return arma::join_rows(v, c);
|
||||
}
|
||||
|
||||
arma::mat Rbm::v_to_h(const arma::mat &visible) const
|
||||
{
|
||||
return visible * m_whv + arma::repmat(m_bhv, visible.n_rows, 1);
|
||||
|
||||
+7
-2
@@ -172,10 +172,15 @@ public:
|
||||
return m_bv.size();
|
||||
}
|
||||
|
||||
private:
|
||||
Params m_params;
|
||||
arma::mat vc_to_v(const arma::mat &vc) const;
|
||||
arma::mat vc_to_c(const arma::mat &vc) const;
|
||||
arma::mat to_vc(const arma::mat &v, const arma::mat &c) const;
|
||||
|
||||
arma::mat v_to_h(const arma::mat &visible) const;
|
||||
arma::mat h_to_v(const arma::mat &hidden) const;
|
||||
|
||||
private:
|
||||
Params m_params;
|
||||
arma::mat sample(arma::mat const &src);
|
||||
void weightUpdate(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv);
|
||||
void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5);
|
||||
|
||||
+63
-38
@@ -26,16 +26,16 @@
|
||||
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, numContext)
|
||||
, m_currWeightIndexToDraw(0)
|
||||
, DrawTraining(nullptr)
|
||||
, DrawReconstruction(nullptr)
|
||||
, DrawVisibleTrain(nullptr)
|
||||
, DrawVisibleReconst(nullptr)
|
||||
, DrawWeights(nullptr)
|
||||
, DrawHidden(nullptr)
|
||||
, DrawContextReconst(nullptr)
|
||||
, m_sizeX(430)
|
||||
, m_sizeY(190)
|
||||
{
|
||||
addAndMakeVisible (DrawTraining = new DrawComponent (numVisibleX, numVisibleY));
|
||||
DrawTraining->setListener(this);
|
||||
addAndMakeVisible (DrawVisibleTrain = new DrawComponent (numVisibleX, numVisibleY));
|
||||
DrawVisibleTrain->setListener(this);
|
||||
|
||||
addAndMakeVisible (DrawHidden = new DrawComponent (numHidden, 1));
|
||||
DrawHidden->setListener(this);
|
||||
@@ -49,7 +49,7 @@ RbmComponent::RbmComponent (const std::string &name, size_t id, size_t numVisibl
|
||||
DrawContextReconst->setListener(this);
|
||||
}
|
||||
|
||||
addAndMakeVisible (DrawReconstruction = new DrawComponent (numVisibleX, numVisibleY));
|
||||
addAndMakeVisible (DrawVisibleReconst = new DrawComponent (numVisibleX, numVisibleY));
|
||||
addAndMakeVisible (DrawWeights = new DrawComponent (numVisibleX, numVisibleY, 0.5, 0.5));
|
||||
|
||||
addAndMakeVisible (m_toggleEnable = new ToggleButton ("Enable layer"));
|
||||
@@ -66,8 +66,8 @@ RbmComponent::RbmComponent (const std::string &name, size_t id, size_t numVisibl
|
||||
|
||||
RbmComponent::~RbmComponent()
|
||||
{
|
||||
DrawTraining = nullptr;
|
||||
DrawReconstruction = nullptr;
|
||||
DrawVisibleTrain = nullptr;
|
||||
DrawVisibleReconst = nullptr;
|
||||
DrawWeights = nullptr;
|
||||
DrawHidden = nullptr;
|
||||
DrawContextTrain = nullptr;
|
||||
@@ -143,8 +143,8 @@ void RbmComponent::paint (Graphics& g)
|
||||
|
||||
void RbmComponent::resized()
|
||||
{
|
||||
DrawTraining->setBoundsRelative(0, 0, 100./m_sizeX, 100./m_sizeY);
|
||||
DrawReconstruction->setBoundsRelative(110./m_sizeX, 0, 100./m_sizeX, 100./m_sizeY);
|
||||
DrawVisibleTrain->setBoundsRelative(0, 0, 100./m_sizeX, 100./m_sizeY);
|
||||
DrawVisibleReconst->setBoundsRelative(110./m_sizeX, 0, 100./m_sizeX, 100./m_sizeY);
|
||||
DrawWeights->setBoundsRelative (220./m_sizeX, 0, 100./m_sizeX, 100./m_sizeY);
|
||||
if (DrawContextTrain)
|
||||
{
|
||||
@@ -214,13 +214,13 @@ void RbmComponent::onDraw(DrawComponent &obj)
|
||||
{
|
||||
downPass(obj.getData());
|
||||
}
|
||||
if (&obj == DrawTraining)
|
||||
if (&obj == DrawVisibleTrain)
|
||||
{
|
||||
upDownPass(obj.getData());
|
||||
upDownPass(getTraining());
|
||||
}
|
||||
if (&obj == DrawContextTrain)
|
||||
{
|
||||
upDownPass(DrawTraining->getData());
|
||||
upDownPass(getTraining());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -252,61 +252,86 @@ void RbmComponent::buttonClicked(Button* buttonThatWasClicked)
|
||||
void RbmComponent::redrawReconstruction()
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (root());
|
||||
pComp->upDownPass(pComp->DrawTraining->getData());
|
||||
pComp->upDownPass(pComp->getTraining());
|
||||
|
||||
}
|
||||
|
||||
void RbmComponent::gibbs(const arma::mat& v)
|
||||
void RbmComponent::gibbs(const arma::mat& vc)
|
||||
{
|
||||
arma::mat r = v;
|
||||
arma::mat c;
|
||||
arma::mat r = vc;
|
||||
for (int i=0; i < params().numGibbs; i++)
|
||||
{
|
||||
DrawHidden->getData() = toHiddenProbs(r);
|
||||
r = toVisibleProbs(DrawHidden->getData());
|
||||
c = toContextProbs(DrawHidden->getData());
|
||||
DrawHidden->getData() = prob(v_to_h(r));
|
||||
r = prob(h_to_v(DrawHidden->getData()));
|
||||
}
|
||||
DrawReconstruction->getData() = r;
|
||||
DrawContextReconst->getData() = c;
|
||||
DrawVisibleReconst->getData() = vc_to_v(r);
|
||||
DrawContextReconst->getData() = vc_to_c(r);
|
||||
|
||||
DrawReconstruction->DrawData();
|
||||
DrawVisibleReconst->DrawData();
|
||||
DrawContextReconst->DrawData();
|
||||
|
||||
DrawHidden->DrawData();
|
||||
DrawContextReconst->DrawData();
|
||||
}
|
||||
|
||||
void RbmComponent::upPass(const arma::mat& v)
|
||||
arma::mat RbmComponent::getTraining() const
|
||||
{
|
||||
DrawTraining->getData() = v;
|
||||
DrawTraining->DrawData();
|
||||
gibbs(v);
|
||||
return to_vc(DrawVisibleTrain->getData(), DrawContextTrain->getData());
|
||||
}
|
||||
|
||||
arma::mat RbmComponent::getReconst() const
|
||||
{
|
||||
return to_vc(DrawVisibleReconst->getData(), DrawContextReconst->getData());
|
||||
}
|
||||
|
||||
void RbmComponent::trainRedraw(const arma::mat& vc)
|
||||
{
|
||||
DrawVisibleTrain->getData() = vc_to_v(vc);
|
||||
DrawContextTrain->getData() = vc_to_c(vc);
|
||||
DrawVisibleTrain->DrawData();
|
||||
DrawContextTrain->DrawData();
|
||||
}
|
||||
|
||||
void RbmComponent::reconstRedraw(const arma::mat& vc)
|
||||
{
|
||||
DrawVisibleReconst->getData() = vc_to_v(vc);
|
||||
DrawContextReconst->getData() = vc_to_c(vc);
|
||||
DrawVisibleReconst->DrawData();
|
||||
DrawContextReconst->DrawData();
|
||||
}
|
||||
|
||||
void RbmComponent::upPass(const arma::mat& vc)
|
||||
{
|
||||
DrawVisibleTrain->getData() = vc_to_v(vc);
|
||||
DrawVisibleTrain->DrawData();
|
||||
DrawContextTrain->getData() = vc_to_c(vc);
|
||||
DrawContextTrain->DrawData();
|
||||
gibbs(vc);
|
||||
if (next)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (next);
|
||||
pComp->upPass(DrawHidden->getData());
|
||||
}
|
||||
pComp->upPass(DrawHidden->getData());
|
||||
}
|
||||
}
|
||||
|
||||
void RbmComponent::downPass(const arma::mat& h)
|
||||
{
|
||||
DrawHidden->getData() = h;
|
||||
DrawHidden->DrawData();
|
||||
DrawReconstruction->getData() = toVisibleProbs(h);
|
||||
DrawReconstruction->DrawData();
|
||||
DrawContextReconst->DrawData();
|
||||
|
||||
reconstRedraw(prob(h_to_v(h)));
|
||||
|
||||
if (prev)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
|
||||
pComp->downPass(DrawReconstruction->getData());
|
||||
pComp->downPass(getReconst());
|
||||
}
|
||||
}
|
||||
|
||||
void RbmComponent::upDownPass(const arma::mat& v)
|
||||
void RbmComponent::upDownPass(const arma::mat& vc)
|
||||
{
|
||||
DrawTraining->getData() = v;
|
||||
DrawTraining->DrawData();
|
||||
gibbs(v);
|
||||
trainRedraw(vc);
|
||||
gibbs(vc);
|
||||
if (next)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (next);
|
||||
@@ -324,7 +349,7 @@ void RbmComponent::upDownPass(const arma::mat& v)
|
||||
if (prev)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
|
||||
pComp->downPass(DrawReconstruction->getData());
|
||||
pComp->downPass(getReconst());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -358,7 +383,7 @@ void RbmComponent::redrawWeights(size_t index)
|
||||
void RbmComponent::setTrainingData(arma::mat const& batch)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*> (root());
|
||||
pComp->upDownPass(batch);
|
||||
pComp->upDownPass(to_vc(batch, pComp->DrawContextTrain->getData()));
|
||||
}
|
||||
//[/MiscUserCode]
|
||||
|
||||
|
||||
@@ -70,14 +70,18 @@ public:
|
||||
void downPass(arma::mat const &v);
|
||||
void upDownPass(arma::mat const &v);
|
||||
arma::mat getConvolutedWeight(arma::mat const &h);
|
||||
ScopedPointer<DrawComponent> DrawTraining;
|
||||
ScopedPointer<DrawComponent> DrawVisibleTrain;
|
||||
ScopedPointer<DrawComponent> DrawHidden;
|
||||
ScopedPointer<DrawComponent> DrawContextTrain;
|
||||
ScopedPointer<DrawComponent> DrawContextReconst;
|
||||
ScopedPointer<DrawComponent> DrawContextTrain;
|
||||
arma::mat getTraining() const;
|
||||
arma::mat getReconst() const;
|
||||
void trainRedraw(const arma::mat& vc);
|
||||
void reconstRedraw(const arma::mat& vc);
|
||||
|
||||
private:
|
||||
//[UserVariables] -- You can add your own custom variables in this section.
|
||||
ScopedPointer<DrawComponent> DrawReconstruction;
|
||||
ScopedPointer<DrawComponent> DrawVisibleReconst;
|
||||
ScopedPointer<DrawComponent> DrawContextReconst;
|
||||
ScopedPointer<DrawComponent> DrawWeights;
|
||||
ScopedPointer<ToggleButton> m_toggleEnable;
|
||||
void gibbs(const arma::mat& v);
|
||||
|
||||
Reference in New Issue
Block a user