more stable
git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@286 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+42
-38
@@ -96,7 +96,7 @@ MainComponent::MainComponent ()
|
||||
testButton->addListener (this);
|
||||
|
||||
addAndMakeVisible (numVisibleLabel = new Label ("Num Visible label",
|
||||
TRANS("99999")));
|
||||
TRANS("16")));
|
||||
numVisibleLabel->setFont (Font (15.00f, Font::plain));
|
||||
numVisibleLabel->setJustificationType (Justification::centred);
|
||||
numVisibleLabel->setEditable (true, true, false);
|
||||
@@ -105,7 +105,7 @@ MainComponent::MainComponent ()
|
||||
numVisibleLabel->addListener (this);
|
||||
|
||||
addAndMakeVisible (numHiddenLabel = new Label ("Num Hidden label",
|
||||
TRANS("99999")));
|
||||
TRANS("64")));
|
||||
numHiddenLabel->setFont (Font (15.00f, Font::plain));
|
||||
numHiddenLabel->setJustificationType (Justification::centred);
|
||||
numHiddenLabel->setEditable (true, true, false);
|
||||
@@ -135,7 +135,7 @@ MainComponent::MainComponent ()
|
||||
saveButton->addListener (this);
|
||||
|
||||
addAndMakeVisible (numVisibleYLabel = new Label ("Num Visible Y label",
|
||||
TRANS("99999")));
|
||||
TRANS("16")));
|
||||
numVisibleYLabel->setFont (Font (15.00f, Font::plain));
|
||||
numVisibleYLabel->setJustificationType (Justification::centred);
|
||||
numVisibleYLabel->setEditable (true, true, false);
|
||||
@@ -273,14 +273,8 @@ MainComponent::MainComponent ()
|
||||
|
||||
|
||||
//[UserPreSize]
|
||||
m_vNumX = 16;
|
||||
m_vNumY = 16;
|
||||
m_hNum = 64;
|
||||
m_numGibbs = 1;
|
||||
m_weights.setUnits(m_vNumX, m_vNumY, m_hNum);
|
||||
m_pRbm = new Rbm(m_weights, this);
|
||||
create();
|
||||
|
||||
//[/UserPreSize]
|
||||
|
||||
setSize (800, 600);
|
||||
@@ -494,7 +488,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
||||
else if (buttonThatWasClicked == ShakeButton)
|
||||
{
|
||||
//[UserButtonCode_ShakeButton] -- add your button handler code here..
|
||||
m_weights.shuffle(weightInitLabel->getText().getFloatValue());
|
||||
m_weights->shuffle(weightInitLabel->getText().getFloatValue());
|
||||
redrawReconstruction();
|
||||
redrawWeights((int)WeightsSlider->getValue());
|
||||
//[/UserButtonCode_ShakeButton]
|
||||
@@ -509,11 +503,6 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
||||
else if (buttonThatWasClicked == createButton)
|
||||
{
|
||||
//[UserButtonCode_createButton] -- add your button handler code here..
|
||||
m_vNumX = numVisibleLabel->getText().getIntValue();
|
||||
m_vNumY = numVisibleYLabel->getText().getIntValue();
|
||||
m_hNum = numHiddenLabel->getText().getIntValue();
|
||||
|
||||
m_weights.setUnits(m_vNumX, m_vNumY, m_hNum);
|
||||
create();
|
||||
//[/UserButtonCode_createButton]
|
||||
}
|
||||
@@ -521,8 +510,6 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
||||
{
|
||||
//[UserButtonCode_loadButton] -- add your button handler code here..
|
||||
load();
|
||||
redrawReconstruction();
|
||||
redrawWeights((int)WeightsSlider->getValue());
|
||||
//[/UserButtonCode_loadButton]
|
||||
}
|
||||
else if (buttonThatWasClicked == saveButton)
|
||||
@@ -536,7 +523,6 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
|
||||
//[UserButtonCode_loadTrainingButton] -- add your button handler code here..
|
||||
m_layers.clear();
|
||||
m_layers.load((String(getBaseDir() + String(".trainingStates.dat"))).toUTF8());
|
||||
redrawReconstruction();
|
||||
//[/UserButtonCode_loadTrainingButton]
|
||||
}
|
||||
else if (buttonThatWasClicked == saveTrainingButton)
|
||||
@@ -799,26 +785,35 @@ void MainComponent::mouseWheelMove (const MouseEvent& e, const MouseWheelDetails
|
||||
//[MiscUserCode] You can add your own definitions of your custom methods or any other code here...
|
||||
void MainComponent::load ()
|
||||
{
|
||||
m_weights.load((String(getBaseDir() + String(".weights.dat"))).toUTF8());
|
||||
|
||||
m_vNumX = m_weights.getNumVisibleX();
|
||||
m_vNumY = m_weights.getNumVisibleY();
|
||||
m_hNum = m_weights.getNumHidden();
|
||||
create();
|
||||
create((String(getBaseDir() + String(".weights.dat"))).toUTF8());
|
||||
}
|
||||
|
||||
void MainComponent::save ()
|
||||
{
|
||||
m_weights.save((String(getBaseDir() + String(".weights.dat"))).toUTF8());
|
||||
m_weights->save((String(getBaseDir() + String(".weights.dat"))).toUTF8());
|
||||
}
|
||||
|
||||
void MainComponent::create()
|
||||
void MainComponent::create(const char *pFilename)
|
||||
{
|
||||
m_weights = nullptr;
|
||||
DrawTraining = nullptr;
|
||||
DrawReconstruction = nullptr;
|
||||
DrawWeights = nullptr;
|
||||
DrawVars = nullptr;
|
||||
DrawHidden = nullptr;
|
||||
m_pRbm = nullptr;
|
||||
|
||||
if (pFilename)
|
||||
{
|
||||
m_weights = new Weights(pFilename);
|
||||
}
|
||||
else
|
||||
{
|
||||
m_weights = new Weights(numVisibleLabel->getText().getIntValue(), numVisibleYLabel->getText().getIntValue(), numHiddenLabel->getText().getIntValue());
|
||||
}
|
||||
m_vNumX = m_weights->getNumVisibleX();
|
||||
m_vNumY = m_weights->getNumVisibleY();
|
||||
m_hNum = m_weights->getNumHidden();
|
||||
|
||||
WeightsSlider->setRange(0, m_hNum-1, 1);
|
||||
|
||||
@@ -833,6 +828,8 @@ void MainComponent::create()
|
||||
numVisibleYLabel->setText(String(m_vNumY), dontSendNotification );
|
||||
numHiddenLabel->setText(String(m_hNum), dontSendNotification );
|
||||
|
||||
m_pRbm = new Rbm(*m_weights, this);
|
||||
|
||||
m_pRbm->setDoRaoBlackwell(rbmDoRaoBlackwellToggleButton->getToggleState());
|
||||
m_pRbm->setUseProbsForHiddenReconstruction(rbmReduceVarianceToggleButton->getToggleState());
|
||||
m_pRbm->setUseVisibleGaussian(rbmUseVisibleGaussianToggleButton->getToggleState());
|
||||
@@ -845,8 +842,10 @@ void MainComponent::create()
|
||||
m_pRbm->setSparsity(sparsityLabel->getText().getFloatValue());
|
||||
m_pRbm->setNumGibbs((uint32_t)numGibbsSlider->getValue());
|
||||
|
||||
|
||||
resized();
|
||||
|
||||
redrawReconstruction();
|
||||
redrawWeights((int)WeightsSlider->getValue());
|
||||
}
|
||||
|
||||
void MainComponent::destroy()
|
||||
@@ -882,25 +881,30 @@ void MainComponent::onDraw(DrawComponent &obj)
|
||||
}
|
||||
if (&obj == DrawTraining)
|
||||
{
|
||||
// DrawReconstruction->setData(m_pRbm->toVisible(m_pRbm->toHidden(obj.getData())));
|
||||
// DrawHidden->setData(m_pRbm->toHidden(obj.getData()));
|
||||
redrawReconstruction();
|
||||
}
|
||||
}
|
||||
|
||||
void MainComponent::redrawReconstruction()
|
||||
{
|
||||
char table[] = " ABCDEFGHIJKLMNOPQRSTUVWXYZ";
|
||||
VectorXd v;
|
||||
MatrixXd m;
|
||||
|
||||
DrawHidden->setData(m_pRbm->toHidden(DrawTraining->getData()));
|
||||
v = m_pRbm->toVisible(DrawHidden->getData());
|
||||
RowVectorXd h = m_pRbm->toHidden(DrawTraining->getData());
|
||||
DrawHidden->setData(h);
|
||||
|
||||
RowVectorXd v = m_pRbm->toVisible(DrawHidden->getData());
|
||||
DrawReconstruction->setData(v);
|
||||
|
||||
// RowVectorXd dataV(m_weights->getNumVisible());
|
||||
// DrawHidden->setData(m_pRbm->toHidden(dataV));
|
||||
|
||||
// RowVectorXd dataH(m_weights->getNumHidden());
|
||||
// VectorXd v = m_pRbm->toVisible(dataH);
|
||||
// DrawReconstruction->setData(v);
|
||||
|
||||
|
||||
if ((m_vNumX == 27) && (m_vNumY == 4))
|
||||
{
|
||||
m = v;
|
||||
char table[] = " ABCDEFGHIJKLMNOPQRSTUVWXYZ";
|
||||
MatrixXd m = m_pRbm->toVisible(DrawHidden->getData());
|
||||
m.resize(m_vNumX, m_vNumY);
|
||||
cout << "Reconstruction" << endl;
|
||||
cout << m << endl;
|
||||
@@ -927,9 +931,9 @@ void MainComponent::redrawReconstruction()
|
||||
|
||||
void MainComponent::redrawWeights(int index)
|
||||
{
|
||||
VectorXd w = m_weights.weights().col(index);
|
||||
VectorXd w = m_weights->weights().col(index);
|
||||
DrawWeights->setData(w);
|
||||
DrawVars->setData(m_weights.sigma());
|
||||
DrawVars->setData(m_weights->sigma());
|
||||
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user