- 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:
+9
-11
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user