- Added gaussian hidden units

This commit is contained in:
2024-01-24 14:47:21 +01:00
parent fb5ac6d124
commit fe9b1951ea
3 changed files with 27 additions and 6 deletions
+4 -2
View File
@@ -217,7 +217,7 @@ MainComponent::MainComponent (const String prjname)
sigmaLabel->addListener (this);
addAndMakeVisible (rbmUseVisibleGaussianToggleButton = new ToggleButton ("rbmUseVisibleGaussian toggle button"));
rbmUseVisibleGaussianToggleButton->setTooltip (TRANS("Unused"));
rbmUseVisibleGaussianToggleButton->setTooltip (TRANS("Use gaussian visible units"));
rbmUseVisibleGaussianToggleButton->setButtonText (TRANS("Use gaussian visible"));
rbmUseVisibleGaussianToggleButton->addListener (this);
@@ -328,7 +328,7 @@ MainComponent::MainComponent (const String prjname)
sizeMiniBatch->addListener (this);
addAndMakeVisible (rbmUseHiddenGaussianToggleButton = new ToggleButton ("rbmUseHiddenGaussian toggle button"));
rbmUseHiddenGaussianToggleButton->setTooltip (TRANS("Unused"));
rbmUseHiddenGaussianToggleButton->setTooltip (TRANS("Use gaussian hidden units"));
rbmUseHiddenGaussianToggleButton->setButtonText (TRANS("Use gaussian hidden"));
rbmUseHiddenGaussianToggleButton->addListener (this);
@@ -644,6 +644,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
else if (buttonThatWasClicked == rbmUseVisibleGaussianToggleButton)
{
//[UserButtonCode_rbmUseVisibleGaussianToggleButton] -- add your button handler code here..
m_pLayer->params().doGaussianVisible = buttonThatWasClicked->getToggleState();
//[/UserButtonCode_rbmUseVisibleGaussianToggleButton]
}
else if (buttonThatWasClicked == rbmDoSparseToggleButton)
@@ -670,6 +671,7 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
else if (buttonThatWasClicked == rbmUseHiddenGaussianToggleButton)
{
//[UserButtonCode_rbmUseHiddenGaussianToggleButton] -- add your button handler code here..
m_pLayer->params().doGaussianHidden = buttonThatWasClicked->getToggleState();
//[/UserButtonCode_rbmUseHiddenGaussianToggleButton]
}
+15 -4
View File
@@ -131,16 +131,23 @@ void Rbm::gibbs_hv(arma::mat &h_probs, arma::mat &v_probs) const
void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::mat &dbhv, arma::mat &dbv)
{
arma::mat v_probs(v_states);
arma::mat h_states = v_to_h(v_states);
arma::mat h_probs = prob(h_states);
arma::mat h_states;
arma::mat h_probs;
// Sample hidden
if (m_params.doRaoBlackwell)
if (m_params.doGaussianVisible)
{
h_probs = v_to_h(v_states);
h_states = h_probs + arma::randn(h_probs.n_rows, h_probs.n_cols);
}
else if (m_params.doRaoBlackwell)
{
h_probs = prob(v_to_h(v_states));
h_states = h_probs;
}
else
{
h_probs = prob(v_to_h(v_states));
h_states = sample(h_probs);
}
@@ -163,7 +170,11 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::
}
// Create hidden representation given v
if (m_params.gibbsDoSampleVisible)
if (m_params.doGaussianHidden)
{
h_probs = v_to_h(v_probs);
}
else if (m_params.gibbsDoSampleVisible)
{
h_probs = prob(v_to_h(sample(v_probs)));
}
+8
View File
@@ -29,6 +29,8 @@ public:
: learningRate(0.1)
, weightDecay(0.0)
, momentum(0.5)
, doGaussianHidden(false)
, doGaussianVisible(false)
, doRaoBlackwell(true)
, gibbsDoSampleVisible(false)
, gibbsDoSampleHidden(true)
@@ -46,6 +48,8 @@ public:
params["weightDecay"] = weightDecay;
params["learningRate"] = learningRate;
params["momentum"] = momentum;
params["doGaussianVisible"] = (int)doGaussianVisible;
params["doGaussianHidden"] = (int)doGaussianHidden;
params["doRaoBlackwell"] = (int)doRaoBlackwell;
params["gibbsDoSampleVisible"] = (int)gibbsDoSampleVisible;
params["gibbsDoSampleHidden"] = (int)gibbsDoSampleHidden;
@@ -63,6 +67,8 @@ public:
weightDecay = params.get("weightDecay", weightDecay).asDouble();
learningRate = params.get("learningRate", learningRate).asDouble();
momentum = params.get("momentum", momentum).asDouble();
doGaussianVisible = params.get("doGaussianVisible", doGaussianVisible) == 1;
doGaussianHidden = params.get("doGaussianHidden", doGaussianHidden) == 1;
doRaoBlackwell = params.get("doRaoBlackwell", doRaoBlackwell) == 1;
gibbsDoSampleVisible = params.get("gibbsDoSampleVisible", gibbsDoSampleVisible) == 1;
gibbsDoSampleHidden = params.get("gibbsDoSampleHidden", gibbsDoSampleHidden) == 1;
@@ -75,6 +81,8 @@ public:
double weightDecay;
double learningRate;
double momentum;
bool doGaussianVisible;
bool doGaussianHidden;
bool doRaoBlackwell;
bool gibbsDoSampleVisible;
bool gibbsDoSampleHidden;