From 2c2039ec2f8c693b1d83d47948e94968ad826a70 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Thu, 17 Oct 2019 19:48:14 +0000 Subject: [PATCH] - added early stop of training - fixed non-visible training data after load git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@554 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- Source/MainComponent.cpp | 18 +++++++++++++----- Source/MainComponent.h | 1 + Source/Rbm.cpp | 11 ++++++++++- Source/Rbm.hpp | 2 +- 4 files changed, 25 insertions(+), 7 deletions(-) diff --git a/Source/MainComponent.cpp b/Source/MainComponent.cpp index d9b90e9..f8a7d4d 100644 --- a/Source/MainComponent.cpp +++ b/Source/MainComponent.cpp @@ -485,15 +485,15 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked) if (buttonThatWasClicked == trainButton) { //[UserButtonCode_trainButton] -- add your button handler code here.. - if (rbmNormalizeDataToggleButton->getToggleState() == true) + if (isThreadRunning()) { - m_trainingData = Rbm::normalizeData(m_layers.data()); + m_doStop = true; } else { - m_trainingData = m_layers.data(); + m_doStop = false; + startThread(); } - startThread(); //[/UserButtonCode_trainButton] } else if (buttonThatWasClicked == addButton) @@ -874,11 +874,19 @@ const juce::String& MainComponent::getBaseDir() void MainComponent::run() { - m_pRbmComponentCurr->train((size_t)numEpochslabel->getText().getIntValue(), (size_t)sizeMiniBatch->getText().getIntValue()); + m_pRbmComponentCurr->train((size_t)numEpochslabel->getText().getIntValue(), (size_t)sizeMiniBatch->getText().getIntValue(), m_doStop); } void MainComponent::onChanged(const LayerArray &obj) { + if (rbmNormalizeDataToggleButton->getToggleState() == true) + { + m_trainingData = Rbm::normalizeData(m_layers.data()); + } + else + { + m_trainingData = m_layers.data(); + } patterSlider->setRange(0, obj.getSize()-1, 1); } diff --git a/Source/MainComponent.h b/Source/MainComponent.h index 7c3b361..8042a0c 100644 --- a/Source/MainComponent.h +++ b/Source/MainComponent.h @@ -90,6 +90,7 @@ private: void run(); String m_baseDir; TooltipWindow m_toolTipWindow; + bool m_doStop; void clearTraining() { m_layers.clear(); diff --git a/Source/Rbm.cpp b/Source/Rbm.cpp index 9e0e77d..f74890e 100644 --- a/Source/Rbm.cpp +++ b/Source/Rbm.cpp @@ -98,7 +98,7 @@ MatrixXd Rbm::normalizeData(MatrixXd const &src) } -void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin) +void Rbm::train(size_t numEpochs, size_t miniBatchSize, bool &doStop) { size_t i; size_t epoch; @@ -124,6 +124,11 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin) m_progress = 0; while (trainingSizeRemain) { + if (doStop) + { + break; + } + cout << "trainingSizeRemain: " << trainingSizeRemain << endl; size_t toSlice = std::min(miniBatchSize, trainingSizeRemain); MatrixXd batch = __batch.block(batchRowIndex, 0, toSlice, m_w.getNumVisible()); @@ -140,6 +145,10 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin) for (epoch=0; epoch < numEpochs; epoch++) { + if (doStop) + { + break; + } onProgressChanged(); // Create hidden layer base on training data diff --git a/Source/Rbm.hpp b/Source/Rbm.hpp index afb2b32..056e769 100644 --- a/Source/Rbm.hpp +++ b/Source/Rbm.hpp @@ -48,7 +48,7 @@ public: void sampleGaussian(MatrixXd &dst, MatrixXd const &src); void sampleGaussian(MatrixXd &srcDst); static MatrixXd normalizeData(MatrixXd const &src); - void train(size_t numEpochs, size_t miniBatchSize, double sigmaMin = 0.05); + void train(size_t numEpochs, size_t miniBatchSize, bool &doStop); double getProgress() const; void toHidden(RowVectorXd &h, RowVectorXd const &v); void toVisible(RowVectorXd &v, RowVectorXd const &h);