- 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_numVisibleX(numVisibleX)
|
||||||
, m_numVisibleY(numVisibleY)
|
, m_numVisibleY(numVisibleY)
|
||||||
, m_numContext(numContext)
|
, m_numContext(numContext)
|
||||||
, m_context(1, numContext)
|
, m_context(0, numContext)
|
||||||
{
|
{
|
||||||
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
|
cout << "Create Layer " << m_name << "." << to_string((int)m_id) << endl;
|
||||||
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
|
m_weightsFile = m_name + "." + to_string((int)m_id) + string(".weights.dat");
|
||||||
|
|||||||
+5
-5
@@ -40,15 +40,15 @@ public:
|
|||||||
{
|
{
|
||||||
if (m_numContext)
|
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);
|
arma::mat ctx = arma::zeros(1, m_numContext);
|
||||||
for (int i=1; i < batch.n_rows; i++)
|
for (int i=1; i < batch.n_rows; i++)
|
||||||
{
|
{
|
||||||
arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx)));
|
arma::mat h = prob(v_to_h(arma::join_rows(batch.row(i-1), ctx)));
|
||||||
ctx = h;
|
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));
|
Rbm::setBatch(trainingData(batch_with_ctx));
|
||||||
}
|
}
|
||||||
else
|
else
|
||||||
@@ -123,9 +123,9 @@ public:
|
|||||||
return vc.submat(0, numVisible() - m_numContext, 0, numVisible() - 1);
|
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:
|
private:
|
||||||
|
|||||||
@@ -682,7 +682,7 @@ void MainComponent::sliderValueChanged (Slider* sliderThatWasMoved)
|
|||||||
{
|
{
|
||||||
//[UserSliderCode_patterSlider] -- add your slider handling code here..
|
//[UserSliderCode_patterSlider] -- add your slider handling code here..
|
||||||
m_trainingIndex = (int)sliderThatWasMoved->getValue();
|
m_trainingIndex = (int)sliderThatWasMoved->getValue();
|
||||||
m_pLayer->setTrainingData(m_stack->trainingData().row(m_trainingIndex));
|
m_pLayer->setTrainingData(trainingAt(m_trainingIndex));
|
||||||
//[/UserSliderCode_patterSlider]
|
//[/UserSliderCode_patterSlider]
|
||||||
}
|
}
|
||||||
else if (sliderThatWasMoved == WeightsSlider)
|
else if (sliderThatWasMoved == WeightsSlider)
|
||||||
@@ -816,7 +816,7 @@ void MainComponent::comboBoxChanged (ComboBox* comboBoxThatHasChanged)
|
|||||||
m_pLayer->redrawWeights(m_weightIndex);
|
m_pLayer->redrawWeights(m_weightIndex);
|
||||||
if (m_stack->trainingData().n_rows > 0)
|
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]
|
//[/UserComboBoxCode_m_rbmSelect]
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -124,6 +124,15 @@ private:
|
|||||||
patterSlider->setRange(0, m_stack->numTraining()-1, 1);
|
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();
|
void updateControls();
|
||||||
bool onProgress(Rbm *pRbm, const Rbm::Status &status) override;
|
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
|
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
|
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)
|
void RbmComponent::trainRedraw(const arma::mat& vc)
|
||||||
@@ -386,7 +386,7 @@ void RbmComponent::redrawWeights(size_t index)
|
|||||||
void RbmComponent::setTrainingData(arma::mat const& batch)
|
void RbmComponent::setTrainingData(arma::mat const& batch)
|
||||||
{
|
{
|
||||||
RbmComponent *pComp = static_cast<RbmComponent*> (root());
|
RbmComponent *pComp = static_cast<RbmComponent*> (root());
|
||||||
pComp->upDownPass(to_vc(batch, pComp->DrawContextTrain->getData()));
|
pComp->upDownPass(batch);
|
||||||
}
|
}
|
||||||
//[/MiscUserCode]
|
//[/MiscUserCode]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user