- 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
This commit is contained in:
@@ -485,15 +485,15 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
|||||||
if (buttonThatWasClicked == trainButton)
|
if (buttonThatWasClicked == trainButton)
|
||||||
{
|
{
|
||||||
//[UserButtonCode_trainButton] -- add your button handler code here..
|
//[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
|
else
|
||||||
{
|
{
|
||||||
m_trainingData = m_layers.data();
|
m_doStop = false;
|
||||||
}
|
|
||||||
startThread();
|
startThread();
|
||||||
|
}
|
||||||
//[/UserButtonCode_trainButton]
|
//[/UserButtonCode_trainButton]
|
||||||
}
|
}
|
||||||
else if (buttonThatWasClicked == addButton)
|
else if (buttonThatWasClicked == addButton)
|
||||||
@@ -874,11 +874,19 @@ const juce::String& MainComponent::getBaseDir()
|
|||||||
|
|
||||||
void MainComponent::run()
|
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)
|
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);
|
patterSlider->setRange(0, obj.getSize()-1, 1);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -90,6 +90,7 @@ private:
|
|||||||
void run();
|
void run();
|
||||||
String m_baseDir;
|
String m_baseDir;
|
||||||
TooltipWindow m_toolTipWindow;
|
TooltipWindow m_toolTipWindow;
|
||||||
|
bool m_doStop;
|
||||||
void clearTraining()
|
void clearTraining()
|
||||||
{
|
{
|
||||||
m_layers.clear();
|
m_layers.clear();
|
||||||
|
|||||||
+10
-1
@@ -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 i;
|
||||||
size_t epoch;
|
size_t epoch;
|
||||||
@@ -124,6 +124,11 @@ void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
|||||||
m_progress = 0;
|
m_progress = 0;
|
||||||
while (trainingSizeRemain)
|
while (trainingSizeRemain)
|
||||||
{
|
{
|
||||||
|
if (doStop)
|
||||||
|
{
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
cout << "trainingSizeRemain: " << trainingSizeRemain << endl;
|
cout << "trainingSizeRemain: " << trainingSizeRemain << endl;
|
||||||
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
|
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
|
||||||
MatrixXd batch = __batch.block(batchRowIndex, 0, toSlice, m_w.getNumVisible());
|
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++)
|
for (epoch=0; epoch < numEpochs; epoch++)
|
||||||
{
|
{
|
||||||
|
if (doStop)
|
||||||
|
{
|
||||||
|
break;
|
||||||
|
}
|
||||||
onProgressChanged();
|
onProgressChanged();
|
||||||
|
|
||||||
// Create hidden layer base on training data
|
// Create hidden layer base on training data
|
||||||
|
|||||||
+1
-1
@@ -48,7 +48,7 @@ public:
|
|||||||
void sampleGaussian(MatrixXd &dst, MatrixXd const &src);
|
void sampleGaussian(MatrixXd &dst, MatrixXd const &src);
|
||||||
void sampleGaussian(MatrixXd &srcDst);
|
void sampleGaussian(MatrixXd &srcDst);
|
||||||
static MatrixXd normalizeData(MatrixXd const &src);
|
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;
|
double getProgress() const;
|
||||||
void toHidden(RowVectorXd &h, RowVectorXd const &v);
|
void toHidden(RowVectorXd &h, RowVectorXd const &v);
|
||||||
void toVisible(RowVectorXd &v, RowVectorXd const &h);
|
void toVisible(RowVectorXd &v, RowVectorXd const &h);
|
||||||
|
|||||||
Reference in New Issue
Block a user