- refactored and cleaned up

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@290 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2016-06-15 20:32:28 +00:00
parent 3d56e4a163
commit 2de8e9ac48
7 changed files with 277 additions and 277 deletions
+24 -14
View File
@@ -36,7 +36,7 @@ RbmComponent::RbmComponent (Weights &weights, RbmComponentListener &listener)
{
//[UserPreSize]
m_pRbm = new Rbm(m_weights, this);
m_pRbm = new Rbm(m_weights, m_layers, this);
m_vNumX = m_weights.getNumVisibleX();
m_vNumY = m_weights.getNumVisibleY();
m_hNum = m_weights.getNumHidden();
@@ -219,7 +219,8 @@ void RbmComponent::onDraw(DrawComponent &obj)
{
if (&obj == DrawHidden)
{
DrawReconstruction->setData(m_pRbm->toVisible(obj.getData()));
m_pRbm->toVisible(DrawReconstruction->getData(), obj.getData());
DrawReconstruction->DrawData();
}
if (&obj == DrawTraining)
{
@@ -230,29 +231,36 @@ void RbmComponent::onDraw(DrawComponent &obj)
void RbmComponent::redrawReconstruction()
{
uint32_t i;
VectorXd V, H;
V = DrawTraining->getData();
for (i=0; i < m_pRbm->getNumGibbs(); i++)
m_pRbm->toHidden(DrawHidden->getData(), DrawTraining->getData());
m_pRbm->toVisible(DrawReconstruction->getData(), DrawHidden->getData());
for (i=0; i < m_pRbm->getNumGibbs()-1; i++)
{
H = m_pRbm->toHidden(V);
DrawHidden->setData(H);
V = m_pRbm->toVisible(H);
DrawReconstruction->setData(V);
m_pRbm->toHidden(DrawHidden->getData(), DrawReconstruction->getData());
m_pRbm->toVisible(DrawReconstruction->getData(), DrawHidden->getData());
}
DrawHidden->DrawData();
DrawReconstruction->DrawData();
}
void RbmComponent::redrawWeights()
{
VectorXd w = m_weights.weights().col(m_currWeightIndexToDraw);
DrawWeights->setData(w);
DrawVars->setData(m_pRbm->getSigma());
DrawWeights->getData() = w;
DrawWeights->DrawData();
DrawVars->getData() = m_pRbm->getVariableSigma();
DrawVars->DrawData();
}
void RbmComponent::setDoLearnVariance(bool value)
{
m_pRbm->setDoLearnVariance(value);
DrawVars->getData() = m_pRbm->getVariableSigma();
DrawVars->DrawData();
redrawReconstruction();
}
void RbmComponent::setDoRaoBlackwell(bool enable)
@@ -304,9 +312,11 @@ void RbmComponent::setNumGibbs(size_t value)
void RbmComponent::setSigma(double value)
{
m_pRbm->setSigma(value);
DrawVars->setData(m_pRbm->getSigma());
redrawReconstruction();
m_pRbm->setConstantSigma(value);
DrawVars->getData() = m_pRbm->getVariableSigma();
DrawVars->DrawData();
redrawReconstruction();
}
//[/MiscUserCode]