- 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:
@@ -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"));
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -65,7 +65,6 @@ private:
|
||||
std::string m_name;
|
||||
Layer *m_pLayers;
|
||||
arma::mat m_trainingBatch;
|
||||
|
||||
};
|
||||
|
||||
|
||||
|
||||
+6
-2
@@ -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);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user