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