[RBM]
- miniBatchSize is parameter of train() git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@307 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
@@ -278,8 +278,8 @@ MainComponent::MainComponent ()
|
|||||||
rbmDoSampleBatch->addListener (this);
|
rbmDoSampleBatch->addListener (this);
|
||||||
|
|
||||||
addAndMakeVisible (sizeMiniBatch = new Label ("sizeMiniBatch",
|
addAndMakeVisible (sizeMiniBatch = new Label ("sizeMiniBatch",
|
||||||
TRANS("0")));
|
TRANS("100")));
|
||||||
sizeMiniBatch->setTooltip (TRANS("Maximum number of training samples per train"));
|
sizeMiniBatch->setTooltip (TRANS("Mini batch size"));
|
||||||
sizeMiniBatch->setFont (Font (15.00f, Font::plain));
|
sizeMiniBatch->setFont (Font (15.00f, Font::plain));
|
||||||
sizeMiniBatch->setJustificationType (Justification::centred);
|
sizeMiniBatch->setJustificationType (Justification::centred);
|
||||||
sizeMiniBatch->setEditable (true, true, false);
|
sizeMiniBatch->setEditable (true, true, false);
|
||||||
@@ -724,7 +724,6 @@ void MainComponent::labelTextChanged (Label* labelThatHasChanged)
|
|||||||
else if (labelThatHasChanged == sizeMiniBatch)
|
else if (labelThatHasChanged == sizeMiniBatch)
|
||||||
{
|
{
|
||||||
//[UserLabelCode_sizeMiniBatch] -- add your label text handling code here..
|
//[UserLabelCode_sizeMiniBatch] -- add your label text handling code here..
|
||||||
m_pRbmComponentCurr->setMiniBatchSize((size_t)labelThatHasChanged->getText().getFloatValue());
|
|
||||||
//[/UserLabelCode_sizeMiniBatch]
|
//[/UserLabelCode_sizeMiniBatch]
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -875,9 +874,7 @@ const juce::String& MainComponent::getBaseDir()
|
|||||||
|
|
||||||
void MainComponent::run()
|
void MainComponent::run()
|
||||||
{
|
{
|
||||||
// trainButton->setEnabled(false);
|
m_pRbmComponentCurr->train((size_t)numEpochslabel->getText().getIntValue(), (size_t)sizeMiniBatch->getText().getIntValue());
|
||||||
m_pRbmComponentCurr->train(numEpochslabel->getText().getIntValue());
|
|
||||||
// trainButton->setEnabled(true);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
void MainComponent::onChanged(const LayerArray &obj)
|
void MainComponent::onChanged(const LayerArray &obj)
|
||||||
@@ -913,7 +910,6 @@ void MainComponent::updateControls()
|
|||||||
numVisibleLabel->setText(String(m_weightsCurr->getNumVisibleX()), dontSendNotification );
|
numVisibleLabel->setText(String(m_weightsCurr->getNumVisibleX()), dontSendNotification );
|
||||||
numVisibleYLabel->setText(String(m_weightsCurr->getNumVisibleY()), dontSendNotification );
|
numVisibleYLabel->setText(String(m_weightsCurr->getNumVisibleY()), dontSendNotification );
|
||||||
numHiddenLabel->setText(String(m_weightsCurr->getNumHidden()), dontSendNotification );
|
numHiddenLabel->setText(String(m_weightsCurr->getNumHidden()), dontSendNotification );
|
||||||
sizeMiniBatch->setText(String(m_pRbmComponentCurr->params().m_miniBatchSize), dontSendNotification);
|
|
||||||
|
|
||||||
WeightsSlider->setRange(0, m_weightsCurr->getNumHidden()-1, 1);
|
WeightsSlider->setRange(0, m_weightsCurr->getNumHidden()-1, 1);
|
||||||
numGibbsSlider->setValue(m_pRbmComponentCurr->params().m_numGibbs);
|
numGibbsSlider->setValue(m_pRbmComponentCurr->params().m_numGibbs);
|
||||||
@@ -1118,8 +1114,8 @@ BEGIN_JUCER_METADATA
|
|||||||
buttonText="Sample training" connectedEdges="0" needsCallback="1"
|
buttonText="Sample training" connectedEdges="0" needsCallback="1"
|
||||||
radioGroupId="0" state="0"/>
|
radioGroupId="0" state="0"/>
|
||||||
<LABEL name="sizeMiniBatch" id="dd9e3e6f2b7b22f8" memberName="sizeMiniBatch"
|
<LABEL name="sizeMiniBatch" id="dd9e3e6f2b7b22f8" memberName="sizeMiniBatch"
|
||||||
virtualName="" explicitFocusOrder="0" pos="1064 20 48 24" tooltip="Maximum number of training samples per train"
|
virtualName="" explicitFocusOrder="0" pos="1064 20 48 24" tooltip="Mini batch size"
|
||||||
edTextCol="ff000000" edBkgCol="0" labelText="0" editableSingleClick="1"
|
edTextCol="ff000000" edBkgCol="0" labelText="100" editableSingleClick="1"
|
||||||
editableDoubleClick="1" focusDiscardsChanges="0" fontname="Default font"
|
editableDoubleClick="1" focusDiscardsChanges="0" fontname="Default font"
|
||||||
fontsize="15" bold="0" italic="0" justification="36"/>
|
fontsize="15" bold="0" italic="0" justification="36"/>
|
||||||
<TOGGLEBUTTON name="rbmUseHiddenGaussian toggle button" id="92e05c920283616b"
|
<TOGGLEBUTTON name="rbmUseHiddenGaussian toggle button" id="92e05c920283616b"
|
||||||
|
|||||||
+9
-15
@@ -17,7 +17,6 @@ Rbm::Rbm(Weights &weights, const MatrixXd &batch)
|
|||||||
, m_variableSigma(weights.getNumVisible())
|
, m_variableSigma(weights.getNumVisible())
|
||||||
, m_progress(0)
|
, m_progress(0)
|
||||||
{
|
{
|
||||||
setMiniBatchSize(batch.rows());
|
|
||||||
Noise_Init(&m_noise, 0x32727155);
|
Noise_Init(&m_noise, 0x32727155);
|
||||||
m_variableSigma.fill(m_params.m_constantSigma);
|
m_variableSigma.fill(m_params.m_constantSigma);
|
||||||
updateHiddenBatch();
|
updateHiddenBatch();
|
||||||
@@ -198,17 +197,17 @@ MatrixXd Rbm::calcZ(MatrixXd &v, MatrixXd &h)
|
|||||||
return t1;
|
return t1;
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::train(uint32_t numEpochs, double sigmaMin)
|
void Rbm::train(size_t numEpochs, size_t miniBatchSize, double sigmaMin)
|
||||||
{
|
{
|
||||||
uint32_t i;
|
size_t i;
|
||||||
uint32_t epoch;
|
size_t epoch;
|
||||||
uint32_t gibbs;
|
size_t gibbs;
|
||||||
|
|
||||||
size_t trainingSize = m_batch.rows();
|
size_t trainingSize = m_batch.rows();
|
||||||
size_t trainingSizeRemain = trainingSize;
|
size_t trainingSizeRemain = trainingSize;
|
||||||
size_t batchRowIndex = 0;
|
size_t batchRowIndex = 0;
|
||||||
|
|
||||||
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(m_params.m_miniBatchSize, trainingSize));
|
double dProgress = 1.0/(numEpochs*(double)trainingSize/std::min(miniBatchSize, trainingSize));
|
||||||
|
|
||||||
MatrixXd dBiasV_curr(MatrixXd::Zero(1, m_w.getNumVisible()));
|
MatrixXd dBiasV_curr(MatrixXd::Zero(1, m_w.getNumVisible()));
|
||||||
MatrixXd dBiasH_curr(MatrixXd::Zero(1, m_w.getNumHidden()));
|
MatrixXd dBiasH_curr(MatrixXd::Zero(1, m_w.getNumHidden()));
|
||||||
@@ -221,14 +220,14 @@ void Rbm::train(uint32_t numEpochs, double sigmaMin)
|
|||||||
while (trainingSizeRemain)
|
while (trainingSizeRemain)
|
||||||
{
|
{
|
||||||
cout << "trainingSizeRemain: " << trainingSizeRemain << endl;
|
cout << "trainingSizeRemain: " << trainingSizeRemain << endl;
|
||||||
size_t toSlice = std::min(m_params.m_miniBatchSize, trainingSizeRemain);
|
size_t toSlice = std::min(miniBatchSize, trainingSizeRemain);
|
||||||
MatrixXd batch = m_batch.block(batchRowIndex, 0, toSlice, m_w.getNumVisible());
|
MatrixXd batch = m_batch.block(batchRowIndex, 0, toSlice, m_w.getNumVisible());
|
||||||
trainingSizeRemain -= toSlice;
|
trainingSizeRemain -= toSlice;
|
||||||
batchRowIndex += toSlice;
|
batchRowIndex += toSlice;
|
||||||
size_t batchSize = batch.rows();
|
size_t batchSize = batch.rows();
|
||||||
double mu_w = m_params.m_muWeights/std::min(m_params.m_miniBatchSize, trainingSize);
|
double mu_w = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||||
double mu_biasV = m_params.m_muWeights/std::min(m_params.m_miniBatchSize, trainingSize);
|
double mu_biasV = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||||
double mu_biasH = m_params.m_muWeights/std::min(m_params.m_miniBatchSize, trainingSize);
|
double mu_biasH = m_params.m_muWeights/std::min(miniBatchSize, trainingSize);
|
||||||
|
|
||||||
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
MatrixXd batch_sampled(batchSize, m_w.getNumVisible());
|
||||||
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
MatrixXd v_sampled(batchSize, m_w.getNumVisible());
|
||||||
@@ -514,11 +513,6 @@ void Rbm::setNumGibbs(size_t value)
|
|||||||
onParamsChanged();
|
onParamsChanged();
|
||||||
}
|
}
|
||||||
|
|
||||||
void Rbm::setMiniBatchSize(size_t size)
|
|
||||||
{
|
|
||||||
m_params.m_miniBatchSize = size;
|
|
||||||
}
|
|
||||||
|
|
||||||
void Rbm::setMuWeights(double value)
|
void Rbm::setMuWeights(double value)
|
||||||
{
|
{
|
||||||
m_params.m_muWeights = value;
|
m_params.m_muWeights = value;
|
||||||
|
|||||||
+1
-4
@@ -37,7 +37,6 @@ public:
|
|||||||
, m_doNormalizeData(false)
|
, m_doNormalizeData(false)
|
||||||
, m_doLearnVariance(false)
|
, m_doLearnVariance(false)
|
||||||
, m_numGibbs(1)
|
, m_numGibbs(1)
|
||||||
, m_miniBatchSize(100)
|
|
||||||
{
|
{
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -58,7 +57,6 @@ public:
|
|||||||
bool m_doNormalizeData;
|
bool m_doNormalizeData;
|
||||||
bool m_doLearnVariance;
|
bool m_doLearnVariance;
|
||||||
size_t m_numGibbs;
|
size_t m_numGibbs;
|
||||||
size_t m_miniBatchSize;
|
|
||||||
};
|
};
|
||||||
|
|
||||||
Rbm(Weights &weights, const MatrixXd &batch);
|
Rbm(Weights &weights, const MatrixXd &batch);
|
||||||
@@ -78,7 +76,7 @@ public:
|
|||||||
RowVectorXd calcMean(MatrixXd const &batch);
|
RowVectorXd calcMean(MatrixXd const &batch);
|
||||||
RowVectorXd calcSigma(MatrixXd const &batch);
|
RowVectorXd calcSigma(MatrixXd const &batch);
|
||||||
MatrixXd calcZ(MatrixXd &v, MatrixXd &h);
|
MatrixXd calcZ(MatrixXd &v, MatrixXd &h);
|
||||||
void train(uint32_t numEpochs, double sigmaMin = 0.05);
|
void train(size_t numEpochs, size_t miniBatchSize, double sigmaMin = 0.05);
|
||||||
double getProgress() const;
|
double getProgress() const;
|
||||||
double getEnergy(const VectorXd& visible, const VectorXd& hidden);
|
double getEnergy(const VectorXd& visible, const VectorXd& hidden);
|
||||||
void toHidden(RowVectorXd &h, RowVectorXd const &v);
|
void toHidden(RowVectorXd &h, RowVectorXd const &v);
|
||||||
@@ -98,7 +96,6 @@ public:
|
|||||||
void setNormalizeData(bool flag);
|
void setNormalizeData(bool flag);
|
||||||
void setDoLearnVariance(bool flag);
|
void setDoLearnVariance(bool flag);
|
||||||
void setNumGibbs(size_t value);
|
void setNumGibbs(size_t value);
|
||||||
void setMiniBatchSize(size_t size);
|
|
||||||
void setMuWeights(double value);
|
void setMuWeights(double value);
|
||||||
void setMuSparsity(double value);
|
void setMuSparsity(double value);
|
||||||
void setMomentum(double value);
|
void setMomentum(double value);
|
||||||
|
|||||||
Reference in New Issue
Block a user