- 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:
2016-06-18 17:57:32 +00:00
parent 46bffda38f
commit 1705287e37
5 changed files with 117 additions and 44 deletions
+23 -8
View File
@@ -508,7 +508,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
//[UserButtonCode_loadButton] -- add your button handler code here..
m_rbmSelect->setSelectedId(1, sendNotification);
create(getBaseDir());
//[/UserButtonCode_loadButton]
//[/UserButtonCode_loadButton]
}
else if (buttonThatWasClicked == saveButton)
{
@@ -520,6 +520,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
{
//[UserButtonCode_loadTrainingButton] -- add your button handler code here..
loadTraining((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
m_pRbmComponent[0]->batchchanged();
//[/UserButtonCode_loadTrainingButton]
}
else if (buttonThatWasClicked == saveTrainingButton)
@@ -532,12 +533,14 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
{
//[UserButtonCode_clearTrainingButton] -- add your button handler code here..
clearTraining();
m_pRbmComponent[0]->batchchanged();
//[/UserButtonCode_clearTrainingButton]
}
else if (buttonThatWasClicked == removeTrainingButton)
{
//[UserButtonCode_removeTrainingButton] -- add your button handler code here..
removeTrainingAt((int)patterSlider->getValue());
m_pRbmComponent[0]->batchchanged();
//[/UserButtonCode_removeTrainingButton]
}
else if (buttonThatWasClicked == reconstructEquButton)
@@ -798,14 +801,18 @@ void MainComponent::save ()
void MainComponent::create(juce::String const &projectName)
{
size_t id = m_rbmSelect->getSelectedId()-1;
if ((m_pRbmComponent[id] == nullptr) and ((id+1) < DBN_SIZE))
m_rbmSelect->addItem(String(id+1), id+2);
m_weights[id] = nullptr;
m_pRbmComponent[id] = nullptr;
bool shouldAddItem = true;
if (id == 0)
{
m_rbmSelect->clear(dontSendNotification);
m_rbmSelect->addItem(String(id), id+1);
m_rbmSelect->setSelectedId(id+1, dontSendNotification);
for (size_t i=0; i < DBN_SIZE; i++)
{
m_weights[i] = nullptr;
m_pRbmComponent[i] = nullptr;
}
if (!projectName.isEmpty())
{
m_weights[id] = new Weights((String(projectName + String(".weights.dat"))).toUTF8());
@@ -819,14 +826,22 @@ void MainComponent::create(juce::String const &projectName)
}
else
{
shouldAddItem = m_pRbmComponent[id] == nullptr;
m_weights[id] = nullptr;
m_pRbmComponent[id] = nullptr;
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));
m_pRbmComponent[id-1]->registerRbm(m_pRbmComponent[id]);
}
m_pRbmComponentCurr = m_pRbmComponent[id];
m_weightsCurr = m_weights[id];
m_pRbmComponentCurr->setBounds (16, 140*id+16, 430, 130);
m_pRbmComponentCurr->batchchanged();
m_pRbmComponentCurr->redrawReconstruction();
if (((id+1) < DBN_SIZE) and shouldAddItem)
m_rbmSelect->addItem(String(id+1), id+2);
updateControls();
}
@@ -854,7 +869,7 @@ void MainComponent::onChanged(const LayerArray &obj)
patterSlider->setRange(0, obj.getSize()-1, 1);
}
void MainComponent::onRbmEpochTrained(size_t progressPercent)
void MainComponent::onProgressChanged(size_t progressPercent)
{
m_progressBarSlider->setValue(progressPercent);
}