[RBM]
- 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:
@@ -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
@@ -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]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
Reference in New Issue
Block a user