- assign training data with context

- calculate context on demand if empty

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@773 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-10 12:19:48 +00:00
parent 9f73824bc4
commit 8c8e0b324d
5 changed files with 21 additions and 12 deletions
+1 -1
View File
@@ -23,7 +23,7 @@ Layer::Layer(const string &name, size_t id, size_t numVisibleX, size_t numVisibl
, m_numVisibleX(numVisibleX)
, m_numVisibleY(numVisibleY)
, m_numContext(numContext)
, m_context(1, numContext)
, m_context(0, numContext)
{
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
+6 -6
View File
@@ -40,15 +40,15 @@ public:
{
if (m_numContext)
{
arma::mat c_states = arma::zeros(batch.n_rows, m_numContext);
m_context = arma::zeros(batch.n_rows, m_numContext);
arma::mat ctx = arma::zeros(1, m_numContext);
for (int i=1; i < batch.n_rows; i++)
{
arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx)));
ctx = h;
c_states.row(i) = ctx;
m_context.row(i) = ctx;
}
arma::mat batch_with_ctx = arma::join_rows(batch, c_states);
arma::mat batch_with_ctx = arma::join_rows(batch, m_context);
Rbm::setBatch(trainingData(batch_with_ctx));
}
else
@@ -123,11 +123,11 @@ public:
return vc.submat(0, numVisible() - m_numContext, 0, numVisible() - 1);
}
arma::mat to_vc(const arma::mat &v, const arma::mat &c) const
const arma::mat& context() const
{
return arma::join_rows(v, c);
return m_context;
}
private:
std::string m_name;
+2 -2
View File
@@ -682,7 +682,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
{
//[UserSliderCode_patterSlider] -- add your slider handling code here..
m_trainingIndex = (int)sliderThatWasMoved->getValue();
m_pLayer->setTrainingData(m_stack->trainingData().row(m_trainingIndex));
m_pLayer->setTrainingData(trainingAt(m_trainingIndex));
//[/UserSliderCode_patterSlider]
}
else if (sliderThatWasMoved == WeightsSlider)
@@ -816,7 +816,7 @@ void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged)
m_pLayer->redrawWeights(m_weightIndex);
if (m_stack->trainingData().n_rows > 0)
{
m_pLayer->setTrainingData(m_stack->trainingData().row(m_trainingIndex));
m_pLayer->setTrainingData(trainingAt(m_trainingIndex));
}
//[/UserComboBoxCode_m_rbmSelect]
}
+9
View File
@@ -124,6 +124,15 @@ private:
patterSlider->setRange(0, m_stack->numTraining()-1, 1);
}
const arma::mat trainingAt(size_t index)
{
if (m_pLayer->context().is_empty())
{
m_pLayer->setBatch(m_stack->trainingData());
}
return arma::join_rows(m_stack->trainingData().row(index), m_pLayer->context().row(index));
}
void updateControls();
bool onProgress(Rbm *pRbm, const Rbm::Status &status) override;
+3 -3
View File
@@ -279,12 +279,12 @@ void RbmComponent::gibbs(const arma::mat& vc)
arma::mat RbmComponent::getTraining() const
{
return to_vc(DrawVisibleTrain->getData(), DrawContextTrain->getData());
return arma::join_rows(DrawVisibleTrain->getData(), DrawContextTrain->getData());
}
arma::mat RbmComponent::getReconst() const
{
return to_vc(DrawVisibleReconst->getData(), DrawContextReconst->getData());
return arma::join_rows(DrawVisibleReconst->getData(), DrawContextReconst->getData());
}
void RbmComponent::trainRedraw(const arma::mat& vc)
@@ -386,7 +386,7 @@ void RbmComponent::redrawWeights(size_t index)
void RbmComponent::setTrainingData(arma::mat const& batch)
{
RbmComponent *pComp = static_cast<RbmComponent*> (root());
pComp->upDownPass(to_vc(batch, pComp->DrawContextTrain->getData()));
pComp->upDownPass(batch);
}
//[/MiscUserCode]