[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:
+82
-7
@@ -23,9 +23,43 @@
|
||||
|
||||
|
||||
//==============================================================================
|
||||
RbmComponent::RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponentListener &listener)
|
||||
: Rbm(weights, batch)
|
||||
, m_weights(weights)
|
||||
RbmComponent::RbmComponent (Weights &_weights, RbmComponent *pRbmUpper, RbmComponentListener &listener)
|
||||
: Rbm(_weights, pRbmUpper->getHiddenBatch())
|
||||
, 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_currWeightIndexToDraw(0)
|
||||
, m_currTrainingIndexToDraw(0)
|
||||
@@ -35,7 +69,6 @@ RbmComponent::RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponen
|
||||
, DrawHidden(nullptr)
|
||||
{
|
||||
|
||||
//[UserPreSize]
|
||||
size_t vNumX = m_weights.getNumVisibleX();
|
||||
size_t vNumY = m_weights.getNumVisibleY();
|
||||
size_t hNum = m_weights.getNumHidden();
|
||||
@@ -44,14 +77,15 @@ RbmComponent::RbmComponent (Weights &weights, MatrixXd const &batch, RbmComponen
|
||||
DrawTraining->setListener(this);
|
||||
|
||||
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 (DrawHidden = new DrawComponent (hNum, 1));
|
||||
DrawHidden->setListener(this);
|
||||
redrawWeights();
|
||||
redrawReconstruction();
|
||||
redrawVariances();
|
||||
//[/UserPreSize]
|
||||
|
||||
setSize (430, 130);
|
||||
resized();
|
||||
@@ -252,9 +286,38 @@ void RbmComponent::redrawReconstruction()
|
||||
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()
|
||||
{
|
||||
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();
|
||||
}
|
||||
|
||||
@@ -315,6 +378,18 @@ RowVectorXd const& RbmComponent::getTrainingData()
|
||||
return DrawTraining->getData();
|
||||
}
|
||||
|
||||
Weights& RbmComponent::getTopWeights()
|
||||
{
|
||||
if (upper)
|
||||
{
|
||||
return upper->getTopWeights();
|
||||
}
|
||||
else
|
||||
{
|
||||
return m_weights;
|
||||
}
|
||||
|
||||
}
|
||||
//[/MiscUserCode]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user