[RBM]
- DBN fixes - batch sample inside training loop - implemented RbmComponent stacking git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@296 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+16
-10
@@ -204,20 +204,20 @@ void RbmComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails&
|
||||
//[/UserCode_mouseWheelMove]
|
||||
}
|
||||
|
||||
void RbmComponent::onEpochTrained(const Rbm &obj)
|
||||
void RbmComponent::onProgressChanged(const Rbm &obj)
|
||||
{
|
||||
// mylog("Training %f %%\n", obj.getProgress()*100);
|
||||
upPass(DrawHidden->getData());
|
||||
redrawReconstruction();
|
||||
redrawWeights();
|
||||
m_listener.onRbmEpochTrained((size_t)(100*obj.getProgress()));
|
||||
m_listener.onProgressChanged((size_t)(100*obj.getProgress() + 0.5));
|
||||
}
|
||||
|
||||
void RbmComponent::onDraw(DrawComponent &obj)
|
||||
{
|
||||
if (&obj == DrawHidden)
|
||||
{
|
||||
m_pRbm->toVisible(DrawReconstruction->getData(), obj.getData());
|
||||
DrawReconstruction->DrawData();
|
||||
upPass(obj.getData());
|
||||
}
|
||||
if (&obj == DrawTraining)
|
||||
{
|
||||
@@ -229,15 +229,16 @@ void RbmComponent::redrawReconstruction()
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
m_pRbm->toHidden(DrawHidden->getData(), DrawTraining->getData());
|
||||
m_pRbm->toVisible(DrawReconstruction->getData(), DrawHidden->getData());
|
||||
reconstructVisible(DrawReconstruction->getData(), DrawTraining->getData());
|
||||
// toHidden(DrawHidden->getData(), DrawTraining->getData());
|
||||
// toVisible(DrawReconstruction->getData(), DrawHidden->getData());
|
||||
|
||||
for (i=0; i < m_pRbm->params().m_numGibbs-1; i++)
|
||||
{
|
||||
m_pRbm->toHidden(DrawHidden->getData(), DrawReconstruction->getData());
|
||||
m_pRbm->toVisible(DrawReconstruction->getData(), DrawHidden->getData());
|
||||
reconstructVisible(DrawReconstruction->getData(), DrawReconstruction->getData());
|
||||
// toHidden(DrawHidden->getData(), DrawReconstruction->getData());
|
||||
// toVisible(DrawReconstruction->getData(), DrawHidden->getData());
|
||||
}
|
||||
DrawHidden->DrawData();
|
||||
DrawReconstruction->DrawData();
|
||||
}
|
||||
|
||||
@@ -337,6 +338,10 @@ void RbmComponent::setMomentum(double value)
|
||||
void RbmComponent::batchchanged()
|
||||
{
|
||||
m_pRbm->updateHiddenBatch();
|
||||
if (lower)
|
||||
{
|
||||
lower->batchchanged();
|
||||
}
|
||||
redrawVariances();
|
||||
}
|
||||
|
||||
@@ -358,8 +363,9 @@ void RbmComponent::setTrainingIndex(size_t index)
|
||||
m_currTrainingIndexToDraw = index;
|
||||
DrawTraining->getData() = m_pRbm->getBatch().row(index);
|
||||
DrawTraining->DrawData();
|
||||
|
||||
|
||||
redrawReconstruction();
|
||||
upPass(DrawHidden->getData());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user