- improved progressIndicator

- fixed Stack::trainingData()
- fixed test::main

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@638 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-08 06:38:09 +00:00
parent 215cfd11fa
commit 9ea6e48fff
5 changed files with 37 additions and 44 deletions
+9 -11
View File
@@ -19,9 +19,7 @@ class RbmListener : public Rbm::IListener
bool onProgress(const Rbm::Status &status)
{
std::cout << "Progress : " << 100*status.progress << " %" << std::endl;
std::cout << "epoch : " << status.epoch << std::endl;
std::cout << "trainingSizeRemain: " << status.trainingSizeRemain << std::endl;
std::cout << "Progress : " << status.progress << " %" << std::endl;
std::cout << "error (per mini batch) = " << status.err << std::endl;
std::cout << "error (total) = " << status.err_total << std::endl;
std::cout << "L1 = " << status.L1 << std::endl;
@@ -84,15 +82,15 @@ int main()
RbmListener statusDisplay;
Stack stack(project);
arma::mat batch = stack.loadTraining();
stack.loadTraining();
printf("Loaded %d training samples\n", (int)batch.n_rows);
printf("Loaded %d training samples\n", (int)stack.trainingData().n_rows);
stack.addTraining(batch, batch.row(1));
printf("Loaded %d training samples\n", (int)batch.n_rows);
stack.addTraining(stack.trainingData().row(1));
printf("Loaded %d training samples\n", (int)stack.trainingData().n_rows);
stack.delTraining(batch, 0);
printf("Loaded %d training samples\n", (int)batch.n_rows);
stack.delTraining(0);
printf("Loaded %d training samples\n", (int)stack.trainingData().n_rows);
#if 1
const int numLayers = 4;
@@ -127,13 +125,13 @@ int main()
#endif
// Train stack
stack.train(batch, &statusDisplay);
stack.train(&statusDisplay);
// Save weights
stack.saveWeights();
Layer *layer = stack.getLayer(0);
arma::mat v = arma::randu(batch.n_rows, layer->bv().n_elem);
arma::mat v = arma::randu(stack.trainingData().n_rows, layer->bv().n_elem);
arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h);
return 0;