- train layer instead of stack
- fixed downPass upPass git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@645 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+33
-1
@@ -36,7 +36,11 @@ public:
|
|||||||
|
|
||||||
arma::mat up_pass(const arma::mat& hidden);
|
arma::mat up_pass(const arma::mat& hidden);
|
||||||
arma::mat down_pass(const arma::mat& visible);
|
arma::mat down_pass(const arma::mat& visible);
|
||||||
|
void train(arma::mat const &batch, IListener *pListener=nullptr)
|
||||||
|
{
|
||||||
|
Rbm::train(trainingData(batch), pListener);
|
||||||
|
}
|
||||||
|
|
||||||
std::string& name()
|
std::string& name()
|
||||||
{
|
{
|
||||||
return m_name;
|
return m_name;
|
||||||
@@ -61,7 +65,35 @@ public:
|
|||||||
{
|
{
|
||||||
return m_bh.n_elem;
|
return m_bh.n_elem;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
arma::mat trainingData(arma::mat const &batch)
|
||||||
|
{
|
||||||
|
arma::mat thisBatch = batch;
|
||||||
|
Layer *pLayer = root();
|
||||||
|
while (pLayer)
|
||||||
|
{
|
||||||
|
if (pLayer == this)
|
||||||
|
{
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
thisBatch = pLayer->toHiddenProbs(thisBatch);
|
||||||
|
pLayer = pLayer->next;
|
||||||
|
}
|
||||||
|
return thisBatch;
|
||||||
|
}
|
||||||
|
|
||||||
|
Layer *root()
|
||||||
|
{
|
||||||
|
Layer *pLayer = this;
|
||||||
|
while(pLayer->prev)
|
||||||
|
{
|
||||||
|
pLayer = pLayer->prev;
|
||||||
|
}
|
||||||
|
return pLayer;
|
||||||
|
}
|
||||||
|
|
||||||
private:
|
private:
|
||||||
|
|
||||||
std::string m_name;
|
std::string m_name;
|
||||||
std::string m_weightsFile;
|
std::string m_weightsFile;
|
||||||
size_t m_id;
|
size_t m_id;
|
||||||
|
|||||||
@@ -835,13 +835,13 @@ const juce::String& MainComponent::getBaseDir()
|
|||||||
|
|
||||||
void MainComponent::run()
|
void MainComponent::run()
|
||||||
{
|
{
|
||||||
m_stack->train(this);
|
m_pLayer->train(m_stack->trainingData(), this);
|
||||||
}
|
}
|
||||||
|
|
||||||
bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
|
bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
|
||||||
{
|
{
|
||||||
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
|
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
|
||||||
pComp->upPass(pComp->DrawHidden->getData());
|
pComp->upPass(pComp->DrawTraining->getData());
|
||||||
pComp->redrawReconstruction();
|
pComp->redrawReconstruction();
|
||||||
pComp->redrawWeights();
|
pComp->redrawWeights();
|
||||||
|
|
||||||
|
|||||||
+33
-29
@@ -207,23 +207,43 @@ void RbmComponent::redrawReconstruction()
|
|||||||
{
|
{
|
||||||
uint32_t i;
|
uint32_t i;
|
||||||
|
|
||||||
downPass(DrawReconstruction->getData(), DrawTraining->getData());
|
for (i=0; i < params().numGibbs; i++)
|
||||||
|
|
||||||
for (i=0; i < params().numGibbs-1; i++)
|
|
||||||
{
|
{
|
||||||
downPass(DrawReconstruction->getData(), DrawReconstruction->getData());
|
downPass(DrawReconstruction->getData());
|
||||||
}
|
}
|
||||||
DrawReconstruction->DrawData();
|
DrawReconstruction->DrawData();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
void RbmComponent::upPass(const arma::mat& v)
|
||||||
|
{
|
||||||
|
DrawReconstruction->getData() = v;
|
||||||
|
DrawReconstruction->DrawData();
|
||||||
|
if (next)
|
||||||
|
{
|
||||||
|
RbmComponent *pComp = static_cast<RbmComponent*> (next);
|
||||||
|
pComp->upPass(toHiddenProbs(v));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
void RbmComponent::downPass(const arma::mat& v)
|
||||||
|
{
|
||||||
|
DrawReconstruction->getData() = v;
|
||||||
|
DrawReconstruction->DrawData();
|
||||||
|
if (prev)
|
||||||
|
{
|
||||||
|
RbmComponent *pComp = static_cast<RbmComponent*> (prev);
|
||||||
|
pComp->downPass(prev->toVisibleProbs(v));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w)
|
arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w)
|
||||||
{
|
{
|
||||||
if (upper)
|
// if (next)
|
||||||
{
|
// {
|
||||||
arma::mat wc = w * upper->getWeights().t(); // * 1.0/sqrt((double)w.cols());
|
// arma::mat wc = w * next->w().t(); // * 1.0/sqrt((double)w.cols());
|
||||||
return upper->getConvolutedWeight(wc);
|
// return static_cast<RbmComponent*>(next)->getConvolutedWeight(wc);
|
||||||
}
|
// }
|
||||||
else
|
// else
|
||||||
{
|
{
|
||||||
return w;
|
return w;
|
||||||
}
|
}
|
||||||
@@ -243,28 +263,12 @@ void RbmComponent::redrawWeights(size_t index)
|
|||||||
|
|
||||||
void RbmComponent::setTrainingData(arma::mat const& batch)
|
void RbmComponent::setTrainingData(arma::mat const& batch)
|
||||||
{
|
{
|
||||||
DrawTraining->getData() = batch;
|
|
||||||
|
DrawTraining->getData() = trainingData(batch);
|
||||||
DrawTraining->DrawData();
|
DrawTraining->DrawData();
|
||||||
|
|
||||||
redrawReconstruction();
|
redrawReconstruction();
|
||||||
upPass(DrawHidden->getData());
|
upPass(DrawTraining->getData());
|
||||||
}
|
|
||||||
|
|
||||||
arma::mat const& RbmComponent::getTopWeights()
|
|
||||||
{
|
|
||||||
if (upper)
|
|
||||||
{
|
|
||||||
return upper->getTopWeights();
|
|
||||||
}
|
|
||||||
else
|
|
||||||
{
|
|
||||||
return w();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
arma::mat const& RbmComponent::getWeights()
|
|
||||||
{
|
|
||||||
return w();
|
|
||||||
}
|
}
|
||||||
//[/MiscUserCode]
|
//[/MiscUserCode]
|
||||||
|
|
||||||
|
|||||||
+3
-51
@@ -26,33 +26,6 @@
|
|||||||
#include "Layer.hpp"
|
#include "Layer.hpp"
|
||||||
//[/Headers]
|
//[/Headers]
|
||||||
|
|
||||||
|
|
||||||
class IRbmComponent
|
|
||||||
{
|
|
||||||
public:
|
|
||||||
IRbmComponent *upper;
|
|
||||||
IRbmComponent *lower;
|
|
||||||
|
|
||||||
IRbmComponent()
|
|
||||||
: upper(nullptr)
|
|
||||||
, lower(nullptr)
|
|
||||||
{}
|
|
||||||
virtual ~IRbmComponent()
|
|
||||||
{
|
|
||||||
}
|
|
||||||
|
|
||||||
virtual void downPass(arma::mat &dst, arma::mat const &src) = 0;
|
|
||||||
virtual void upPass(arma::mat const &v) = 0;
|
|
||||||
void registerRbm(IRbmComponent *pObj)
|
|
||||||
{
|
|
||||||
lower = pObj;
|
|
||||||
pObj->upper = this;
|
|
||||||
}
|
|
||||||
virtual arma::mat const& getTopWeights() = 0;
|
|
||||||
virtual arma::mat getConvolutedWeight(arma::mat const &h) = 0;
|
|
||||||
virtual arma::mat const& getWeights() = 0;
|
|
||||||
};
|
|
||||||
|
|
||||||
//==============================================================================
|
//==============================================================================
|
||||||
/**
|
/**
|
||||||
//[Comments]
|
//[Comments]
|
||||||
@@ -64,7 +37,6 @@ public:
|
|||||||
class RbmComponent : public Component
|
class RbmComponent : public Component
|
||||||
, public Layer
|
, public Layer
|
||||||
, public DrawListener
|
, public DrawListener
|
||||||
, public IRbmComponent
|
|
||||||
{
|
{
|
||||||
public:
|
public:
|
||||||
//==============================================================================
|
//==============================================================================
|
||||||
@@ -93,30 +65,10 @@ public:
|
|||||||
|
|
||||||
void redrawReconstruction();
|
void redrawReconstruction();
|
||||||
|
|
||||||
void downPass(arma::mat &dst, arma::mat const &src) override
|
void downPass(arma::mat const &v);
|
||||||
{
|
void upPass(arma::mat const &v);
|
||||||
DrawHidden->getData() = toHiddenProbs(src);
|
|
||||||
if (lower)
|
|
||||||
{
|
|
||||||
lower->downPass(DrawHidden->getData(), DrawHidden->getData());
|
|
||||||
}
|
|
||||||
dst = toVisibleProbs(DrawHidden->getData());
|
|
||||||
DrawHidden->DrawData();
|
|
||||||
}
|
|
||||||
|
|
||||||
void upPass(arma::mat const &h) override
|
|
||||||
{
|
|
||||||
DrawReconstruction->getData() = toVisibleProbs(h);
|
|
||||||
DrawReconstruction->DrawData();
|
|
||||||
if (upper)
|
|
||||||
{
|
|
||||||
upper->upPass(DrawReconstruction->getData());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
arma::mat const& getTopWeights() override;
|
arma::mat getConvolutedWeight(arma::mat const &h);
|
||||||
arma::mat getConvolutedWeight(arma::mat const &h) override;
|
|
||||||
arma::mat const& getWeights() override;
|
|
||||||
ScopedPointer<DrawComponent> DrawTraining;
|
ScopedPointer<DrawComponent> DrawTraining;
|
||||||
ScopedPointer<DrawComponent> DrawHidden;
|
ScopedPointer<DrawComponent> DrawHidden;
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user