- render lower layer weights as linear combination from top weights using up pass of reconstructions

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@303 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-06-29 20:29:17 +00:00
parent 54e4065a00
commit 6e93c07863
4 changed files with 93 additions and 11 deletions
+3 -2
View File
@@ -815,14 +815,15 @@ void MainComponent::create(juce::String const &projectName)
m_weights[id] = nullptr; m_weights[id] = nullptr;
m_pRbmComponent[id] = nullptr; m_pRbmComponent[id] = nullptr;
m_weights[id] = new Weights(m_weights[id-1]->getNumHidden(), 1, numHiddenLabel->getText().getIntValue()); m_weights[id] = new Weights(m_weights[id-1]->getNumHidden(), 1, numHiddenLabel->getText().getIntValue());
addAndMakeVisible(m_pRbmComponent[id] = new RbmComponent(*m_weights[id], m_pRbmComponent[id-1]->getHiddenBatch(), *this)); addAndMakeVisible(m_pRbmComponent[id] = new RbmComponent(*m_weights[id], m_pRbmComponent[id-1], *this));
m_pRbmComponent[id-1]->registerRbm(m_pRbmComponent[id]);
} }
m_pRbmComponentCurr = m_pRbmComponent[id]; m_pRbmComponentCurr = m_pRbmComponent[id];
m_weightsCurr = m_weights[id]; m_weightsCurr = m_weights[id];
m_pRbmComponentCurr->setBounds (16, 140*id+16, 430, 130); m_pRbmComponentCurr->setBounds (16, 140*id+16, 430, 130);
m_pRbmComponentCurr->batchchanged(); m_pRbmComponentCurr->batchchanged();
m_pRbmComponentCurr->redrawReconstruction(); m_pRbmComponentCurr->redrawReconstruction();
m_pRbmComponentCurr->redrawWeights();
m_pRbmComponentCurr->redrawVariances();
if (((id+1) < DBN_SIZE) and shouldAddItem) if (((id+1) < DBN_SIZE) and shouldAddItem)
m_rbmSelect->addItem(String(id+1), id+2); m_rbmSelect->addItem(String(id+1), id+2);
+82 -7
View File
@@ -23,9 +23,43 @@
//============================================================================== //==============================================================================
RbmComponent::RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponentListener &listener) RbmComponent::RbmComponent (Weights &_weights, RbmComponent *pRbmUpper, RbmComponentListener &listener)
: Rbm(weights, batch) : Rbm(_weights, pRbmUpper->getHiddenBatch())
, m_weights(weights) , m_weights(_weights)
, m_listener(listener)
, m_currWeightIndexToDraw(0)
, m_currTrainingIndexToDraw(0)
, DrawTraining(nullptr)
, DrawReconstruction(nullptr)
, DrawWeights(nullptr)
, DrawHidden(nullptr)
{
pRbmUpper->registerRbm(this);
size_t vNumX = m_weights.getNumVisibleX();
size_t vNumY = m_weights.getNumVisibleY();
size_t hNum = m_weights.getNumHidden();
addAndMakeVisible (DrawTraining = new DrawComponent (vNumX, vNumY));
DrawTraining->setListener(this);
addAndMakeVisible (DrawReconstruction = new DrawComponent (vNumX, vNumY));
addAndMakeVisible (DrawWeights = new DrawComponent (getTopWeights().getNumVisibleX(), getTopWeights().getNumVisibleY(), 0.5, 0.5));
addAndMakeVisible (DrawVars = new DrawComponent (vNumX, vNumY));
addAndMakeVisible (DrawHidden = new DrawComponent (hNum, 1));
DrawHidden->setListener(this);
redrawWeights();
redrawReconstruction();
redrawVariances();
setSize (430, 130);
resized();
}
RbmComponent::RbmComponent (Weights &_weights, MatrixXd const &batch, RbmComponentListener &listener)
: Rbm(_weights, batch)
, m_weights(_weights)
, m_listener(listener) , m_listener(listener)
, m_currWeightIndexToDraw(0) , m_currWeightIndexToDraw(0)
, m_currTrainingIndexToDraw(0) , m_currTrainingIndexToDraw(0)
@@ -35,7 +69,6 @@ RbmComponent::RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponen
, DrawHidden(nullptr) , DrawHidden(nullptr)
{ {
//[UserPreSize]
size_t vNumX = m_weights.getNumVisibleX(); size_t vNumX = m_weights.getNumVisibleX();
size_t vNumY = m_weights.getNumVisibleY(); size_t vNumY = m_weights.getNumVisibleY();
size_t hNum = m_weights.getNumHidden(); size_t hNum = m_weights.getNumHidden();
@@ -44,14 +77,15 @@ RbmComponent::RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponen
DrawTraining->setListener(this); DrawTraining->setListener(this);
addAndMakeVisible (DrawReconstruction = new DrawComponent (vNumX, vNumY)); addAndMakeVisible (DrawReconstruction = new DrawComponent (vNumX, vNumY));
addAndMakeVisible (DrawWeights = new DrawComponent (vNumX, vNumY, 0.5, 0.5));
addAndMakeVisible (DrawWeights = new DrawComponent (getTopWeights().getNumVisibleX(), getTopWeights().getNumVisibleY(), 0.5, 0.5));
addAndMakeVisible (DrawVars = new DrawComponent (vNumX, vNumY)); addAndMakeVisible (DrawVars = new DrawComponent (vNumX, vNumY));
addAndMakeVisible (DrawHidden = new DrawComponent (hNum, 1)); addAndMakeVisible (DrawHidden = new DrawComponent (hNum, 1));
DrawHidden->setListener(this); DrawHidden->setListener(this);
redrawWeights(); redrawWeights();
redrawReconstruction(); redrawReconstruction();
redrawVariances(); redrawVariances();
//[/UserPreSize]
setSize (430, 130); setSize (430, 130);
resized(); resized();
@@ -252,9 +286,38 @@ void RbmComponent::redrawReconstruction()
DrawReconstruction->DrawData(); DrawReconstruction->DrawData();
} }
MatrixXd RbmComponent::getAccumulatedWeight(RowVectorXd const &h)
{
if (upper)
{
RowVectorXd v(m_weights.getNumVisible());
toVisible(v, h);
return upper->getAccumulatedWeight(v);
}
else
{
MatrixXd w = h * m_weights.weights().transpose();
// cout << "w=" << endl;
// cout << w << endl;
return w;
}
}
void RbmComponent::redrawWeights() void RbmComponent::redrawWeights()
{ {
DrawWeights->getData() = m_weights.weights().col(m_currWeightIndexToDraw); if (upper)
{
RowVectorXd h = RowVectorXd::Zero(m_weights.getNumHidden());
h(m_currWeightIndexToDraw) = 1.0;
DrawWeights->getData() = getAccumulatedWeight(h);
}
else
{
RowVectorXd w = m_weights.weights().col(m_currWeightIndexToDraw);
// cout << "w=" << endl;
// cout << w << endl;
DrawWeights->getData() = w;
}
DrawWeights->DrawData(); DrawWeights->DrawData();
} }
@@ -315,6 +378,18 @@ RowVectorXd const& RbmComponent::getTrainingData()
return DrawTraining->getData(); return DrawTraining->getData();
} }
Weights& RbmComponent::getTopWeights()
{
if (upper)
{
return upper->getTopWeights();
}
else
{
return m_weights;
}
}
//[/MiscUserCode] //[/MiscUserCode]
+7 -1
View File
@@ -49,6 +49,8 @@ public:
lower = pObj; lower = pObj;
pObj->upper = this; pObj->upper = this;
} }
virtual Weights& getTopWeights() = 0;
virtual MatrixXd getAccumulatedWeight(RowVectorXd const &h) = 0;
}; };
class RbmComponentListener class RbmComponentListener
@@ -77,7 +79,8 @@ class RbmComponent : public Component
{ {
public: public:
//============================================================================== //==============================================================================
RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponentListener &listener); RbmComponent (Weights &_weights, MatrixXd const &batch, RbmComponentListener &listener);
RbmComponent (Weights &_weights, RbmComponent *pRbmUpper, RbmComponentListener &listener);
~RbmComponent(); ~RbmComponent();
//============================================================================== //==============================================================================
@@ -128,6 +131,9 @@ public:
} }
} }
Weights& getTopWeights() override;
MatrixXd getAccumulatedWeight(RowVectorXd const &h) override;
private: private:
//[UserVariables] -- You can add your own custom variables in this section. //[UserVariables] -- You can add your own custom variables in this section.
Weights &m_weights; Weights &m_weights;