- improved multi layer weight reconstruction using weight convolution

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@306 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-07-07 22:54:54 +00:00
parent 57cdfd7065
commit 0259727a2c
2 changed files with 13 additions and 23 deletions
+9 -21
View File
@@ -286,38 +286,22 @@ void RbmComponent::redrawReconstruction()
DrawReconstruction->DrawData();
}
MatrixXd RbmComponent::getAccumulatedWeight(RowVectorXd const &h)
MatrixXd RbmComponent::getConvolutedWeight(RowVectorXd const &w)
{
if (upper)
{
RowVectorXd v(m_weights.getNumVisible());
toVisible(v, h);
return upper->getAccumulatedWeight(v);
MatrixXd wc = w * upper->getWeights().transpose() * 1.0/sqrt((double)w.cols());
return upper->getConvolutedWeight(wc);
}
else
{
MatrixXd w = h * m_weights.weights().transpose();
// cout << "w=" << endl;
// cout << w << endl;
return w;
}
}
void RbmComponent::redrawWeights()
{
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->getData() = getConvolutedWeight(m_weights.weights().col(m_currWeightIndexToDraw));
DrawWeights->DrawData();
}
@@ -388,7 +372,11 @@ Weights& RbmComponent::getTopWeights()
{
return m_weights;
}
}
MatrixXd const& RbmComponent::getWeights()
{
return m_weights.weights();
}
//[/MiscUserCode]