- 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 down_pass(const arma::mat& visible);
|
||||
|
||||
void train(arma::mat const &batch, IListener *pListener=nullptr)
|
||||
{
|
||||
Rbm::train(trainingData(batch), pListener);
|
||||
}
|
||||
|
||||
std::string& name()
|
||||
{
|
||||
return m_name;
|
||||
@@ -61,7 +65,35 @@ public:
|
||||
{
|
||||
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:
|
||||
|
||||
std::string m_name;
|
||||
std::string m_weightsFile;
|
||||
size_t m_id;
|
||||
|
||||
@@ -835,13 +835,13 @@ const juce::String& MainComponent::getBaseDir()
|
||||
|
||||
void MainComponent::run()
|
||||
{
|
||||
m_stack->train(this);
|
||||
m_pLayer->train(m_stack->trainingData(), this);
|
||||
}
|
||||
|
||||
bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
|
||||
{
|
||||
RbmComponent *pComp = static_cast<RbmComponent*>(pRbm);
|
||||
pComp->upPass(pComp->DrawHidden->getData());
|
||||
pComp->upPass(pComp->DrawTraining->getData());
|
||||
pComp->redrawReconstruction();
|
||||
pComp->redrawWeights();
|
||||
|
||||
|
||||
+33
-29
@@ -207,23 +207,43 @@ void RbmComponent::redrawReconstruction()
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
downPass(DrawReconstruction->getData(), DrawTraining->getData());
|
||||
|
||||
for (i=0; i < params().numGibbs-1; i++)
|
||||
for (i=0; i < params().numGibbs; i++)
|
||||
{
|
||||
downPass(DrawReconstruction->getData(), DrawReconstruction->getData());
|
||||
downPass(DrawReconstruction->getData());
|
||||
}
|
||||
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)
|
||||
{
|
||||
if (upper)
|
||||
{
|
||||
arma::mat wc = w * upper->getWeights().t(); // * 1.0/sqrt((double)w.cols());
|
||||
return upper->getConvolutedWeight(wc);
|
||||
}
|
||||
else
|
||||
// if (next)
|
||||
// {
|
||||
// arma::mat wc = w * next->w().t(); // * 1.0/sqrt((double)w.cols());
|
||||
// return static_cast<RbmComponent*>(next)->getConvolutedWeight(wc);
|
||||
// }
|
||||
// else
|
||||
{
|
||||
return w;
|
||||
}
|
||||
@@ -243,28 +263,12 @@ void RbmComponent::redrawWeights(size_t index)
|
||||
|
||||
void RbmComponent::setTrainingData(arma::mat const& batch)
|
||||
{
|
||||
DrawTraining->getData() = batch;
|
||||
|
||||
DrawTraining->getData() = trainingData(batch);
|
||||
DrawTraining->DrawData();
|
||||
|
||||
redrawReconstruction();
|
||||
upPass(DrawHidden->getData());
|
||||
}
|
||||
|
||||
arma::mat const& RbmComponent::getTopWeights()
|
||||
{
|
||||
if (upper)
|
||||
{
|
||||
return upper->getTopWeights();
|
||||
}
|
||||
else
|
||||
{
|
||||
return w();
|
||||
}
|
||||
}
|
||||
|
||||
arma::mat const& RbmComponent::getWeights()
|
||||
{
|
||||
return w();
|
||||
upPass(DrawTraining->getData());
|
||||
}
|
||||
//[/MiscUserCode]
|
||||
|
||||
|
||||
+3
-51
@@ -26,33 +26,6 @@
|
||||
#include "Layer.hpp"
|
||||
//[/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]
|
||||
@@ -64,7 +37,6 @@ public:
|
||||
class RbmComponent : public Component
|
||||
, public Layer
|
||||
, public DrawListener
|
||||
, public IRbmComponent
|
||||
{
|
||||
public:
|
||||
//==============================================================================
|
||||
@@ -93,30 +65,10 @@ public:
|
||||
|
||||
void redrawReconstruction();
|
||||
|
||||
void downPass(arma::mat &dst, arma::mat const &src) override
|
||||
{
|
||||
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());
|
||||
}
|
||||
}
|
||||
void downPass(arma::mat const &v);
|
||||
void upPass(arma::mat const &v);
|
||||
|
||||
arma::mat const& getTopWeights() override;
|
||||
arma::mat getConvolutedWeight(arma::mat const &h) override;
|
||||
arma::mat const& getWeights() override;
|
||||
arma::mat getConvolutedWeight(arma::mat const &h);
|
||||
ScopedPointer<DrawComponent> DrawTraining;
|
||||
ScopedPointer<DrawComponent> DrawHidden;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user