- 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)
{
//[UserButtonCode_addButton] -- add your button handler code here..
Layer *layer = 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]
}
else if (buttonThatWasClicked == ShakeButton)
@@ -899,8 +901,8 @@ const juce::String& MainComponent::getBaseDir()
void MainComponent::run()
{
trainButton->setButtonText (TRANS("Stop"));
m_pLayer->setBatch(m_stack->trainingBatch());
m_pLayer->train(this);
m_pLayer->calcContextBatch(m_stack->trainingBatch());
m_pLayer->train(m_stack->trainingBatch(), this);
trainButton->setButtonText (TRANS("Train"));
}
+1 -12
View File
@@ -126,18 +126,7 @@ private:
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();
+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);
}
void Rbm::train(IListener* pListener)
void Rbm::train(arma::mat const &batch, IListener* pListener)
{
Status status;
arma::mat const &batch = m_batch;
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
double progress = 0;
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;
}
void setBatch(arma::mat const &batch)
{
m_batch = batch;
}
void train(IListener *pListener=nullptr);
void train(arma::mat const &batch, IListener *pListener=nullptr);
static arma::mat normalize(const arma::mat &hidden);
const arma::mat& whv() const;
@@ -196,7 +191,6 @@ protected:
private:
noise_gen_t m_noise;
arma::mat m_whv;
arma::mat m_batch;
};
+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;
}
-1
View File
@@ -65,7 +65,6 @@ private:
std::string m_name;
Layer *m_pLayers;
arma::mat m_trainingBatch;
};
+6 -2
View File
@@ -81,8 +81,6 @@ int main()
Stack stack(".", project);
stack.loadTrainingBatch();
#if CREATE_TEST
stack.addTraining(stack.trainingBatch().row(1));
printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows);
@@ -119,6 +117,10 @@ int main()
// Load weights
stack.loadWeights();
// Load training
stack.loadTrainingBatch();
#endif
#if TRAIN_TEST
@@ -132,6 +134,8 @@ int main()
#endif
Layer *layer = stack.getLayer(0);
arma::mat v = stack.trainingBatch();
v.print("t");
layer->calcContextBatch(v);
arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h);