- 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:
2019-11-08 18:23:56 +00:00
parent 8d9b16c634
commit 9b1084bda9
4 changed files with 71 additions and 83 deletions
+33 -1
View File
@@ -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;
+2 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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;