- 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
+5 -3
View File
@@ -532,8 +532,10 @@ void MainComponent::buttonClicked (Button* buttonThatWasClicked)
else if (buttonThatWasClicked == addButton) else if (buttonThatWasClicked == addButton)
{ {
//[UserButtonCode_addButton] -- add your button handler code here.. //[UserButtonCode_addButton] -- add your button handler code here..
Layer *layer = m_stack->getLayer(0);
RbmComponent *pComp = static_cast<RbmComponent*>(m_stack->getLayer(0)); RbmComponent *pComp = static_cast<RbmComponent*>(m_stack->getLayer(0));
addTraining(pComp->DrawVisibleTrain->getData()); addTraining(arma::join_rows(pComp->DrawVisibleTrain->getData(), pComp->DrawContextTrain->getData()));
//[/UserButtonCode_addButton] //[/UserButtonCode_addButton]
} }
else if (buttonThatWasClicked == ShakeButton) else if (buttonThatWasClicked == ShakeButton)
@@ -899,8 +901,8 @@ const juce::String& MainComponent::getBaseDir()
void MainComponent::run() void MainComponent::run()
{ {
trainButton->setButtonText (TRANS("Stop")); trainButton->setButtonText (TRANS("Stop"));
m_pLayer->setBatch(m_stack->trainingBatch()); m_pLayer->calcContextBatch(m_stack->trainingBatch());
m_pLayer->train(this); m_pLayer->train(m_stack->trainingBatch(), this);
trainButton->setButtonText (TRANS("Train")); trainButton->setButtonText (TRANS("Train"));
} }
-11
View File
@@ -125,20 +125,9 @@ private:
} }
const arma::mat trainingAt(size_t index) const arma::mat trainingAt(size_t index)
{
if (m_pLayer->context().n_cols > 0)
{
if (m_pLayer->context().n_rows != m_stack->trainingBatch().n_rows)
{
m_pLayer->setBatch(m_stack->trainingBatch());
}
return arma::join_rows(m_stack->trainingBatch().row(index), m_pLayer->context().row(index));
}
else
{ {
return m_stack->trainingBatch().row(index); return m_stack->trainingBatch().row(index);
} }
}
void updateControls(); void updateControls();
bool onProgress(Rbm *pRbm, const Rbm::Status &status) override; bool onProgress(Rbm *pRbm, const Rbm::Status &status) override;
+1 -2
View File
@@ -116,10 +116,9 @@ void Rbm::contrastiveDivergence(arma::mat const &v_states, arma::mat &dw, arma::
dbhv -= sum(h_probs, 0); dbhv -= sum(h_probs, 0);
} }
void Rbm::train(IListener* pListener) void Rbm::train(arma::mat const &batch, IListener* pListener)
{ {
Status status; Status status;
arma::mat const &batch = m_batch;
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs); double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
double progress = 0; double progress = 0;
int lastProgress = -100; int lastProgress = -100;
+1 -7
View File
@@ -126,12 +126,7 @@ public:
m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv; m_bv.submat(0, 0, bv.n_rows-1, bv.n_cols-1) = bv;
} }
void setBatch(arma::mat const &batch) void train(arma::mat const &batch, IListener *pListener=nullptr);
{
m_batch = batch;
}
void train(IListener *pListener=nullptr);
static arma::mat normalize(const arma::mat &hidden); static arma::mat normalize(const arma::mat &hidden);
const arma::mat& whv() const; const arma::mat& whv() const;
@@ -196,7 +191,6 @@ protected:
private: private:
noise_gen_t m_noise; noise_gen_t m_noise;
arma::mat m_whv; arma::mat m_whv;
arma::mat m_batch;
}; };
+12 -2
View File
@@ -204,8 +204,8 @@ void Stack::train(Rbm::IListener* pListener)
while(pLayer) while(pLayer)
{ {
std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl; std::cout << m_name << ": " << " Training of layer " << std::to_string(pLayer->id()) << std::endl;
pLayer->setBatch(m_trainingBatch); pLayer->calcContextBatch(m_trainingBatch);
pLayer->train(pListener); pLayer->train(m_trainingBatch, pListener);
pLayer = pLayer->next; pLayer = pLayer->next;
} }
} }
@@ -245,6 +245,16 @@ size_t Stack::loadTrainingBatch(bool doNormalize)
} }
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; 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; return m_trainingBatch.n_rows;
} }
-1
View File
@@ -65,7 +65,6 @@ private:
std::string m_name; std::string m_name;
Layer *m_pLayers; Layer *m_pLayers;
arma::mat m_trainingBatch; arma::mat m_trainingBatch;
}; };
+6 -2
View File
@@ -81,8 +81,6 @@ int main()
Stack stack(".", project); Stack stack(".", project);
stack.loadTrainingBatch();
#if CREATE_TEST #if CREATE_TEST
stack.addTraining(stack.trainingBatch().row(1)); stack.addTraining(stack.trainingBatch().row(1));
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
@@ -119,6 +117,10 @@ int main()
// Load weights // Load weights
stack.loadWeights(); stack.loadWeights();
// Load training
stack.loadTrainingBatch();
#endif #endif
#if TRAIN_TEST #if TRAIN_TEST
@@ -132,6 +134,8 @@ int main()
#endif #endif
Layer *layer = stack.getLayer(0); Layer *layer = stack.getLayer(0);
arma::mat v = stack.trainingBatch(); arma::mat v = stack.trainingBatch();
v.print("t");
layer->calcContextBatch(v);
arma::mat h = layer->toHiddenProbs(v); arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h); arma::mat r = layer->toVisibleProbs(h);