- trainingdata always contains context

- load / store training batch with context
- on load: add context part to  legacy training batches 
- removed Rbm::setBatch()

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@790 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-12 08:12:35 +00:00
parent c88ab5336b
commit fa2014b4dd
7 changed files with 26 additions and 29 deletions
+12 -2
View File
@@ -204,8 +204,8 @@ void Stack::train(Rbm::IListener* pListener)
while(pLayer)
{
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
pLayer->setBatch(m_trainingBatch);
pLayer->train(pListener);
pLayer->calcContextBatch(m_trainingBatch);
pLayer->train(m_trainingBatch, pListener);
pLayer = pLayer->next;
}
}
@@ -245,6 +245,16 @@ size_t Stack::loadTrainingBatch(bool doNormalize)
}
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
}
// Migrate context part to training data
size_t numContext = getLayer(0)->context().n_cols;
size_t numTraining = m_trainingBatch.n_rows;
if (numContext > 0 and m_trainingBatch.n_cols != getLayer(0)->numVisible())
{
arma::mat training_with_ctx = arma::join_rows(m_trainingBatch, arma::zeros(numTraining, numContext));
m_trainingBatch = training_with_ctx;
}
return m_trainingBatch.n_rows;
}