- 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:
+1
-1
@@ -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
@@ -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;
|
||||
|
||||
@@ -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]
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user