- 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 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;
+2 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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;