From b69c735f5db6e1c476cb536644a8036346e0529e Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Wed, 12 Jan 2022 18:18:43 +0000 Subject: [PATCH] - create training data git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@795 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/poet.cpp | 117 ++++++++++++++++++++---------------------------- 1 file changed, 49 insertions(+), 68 deletions(-) diff --git a/source/poet.cpp b/source/poet.cpp index e93e780..bbb69c5 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -29,79 +29,60 @@ class RbmListener : public Rbm::IListener } }; +arma::mat createTraining(const string &filename) +{ + FILE *pFile; + + pFile = fopen(filename.c_str(), "r"); + + if (!pFile) + { + std::cout << "Could not open " << filename << "!" << std::endl; + return 0; + } + + int numTraining = 0; + int numVisible = 26 + 10; + + arma::mat batch = arma::zeros(numTraining, numVisible); + while(!feof(pFile)) + { + char c; + int result = fread(&c, 1, 1, pFile); + if (result < 0) + { + break; + } + + c = toupper(c); + + int index; + if (isalpha(c)) + { + index = (int)(c-'A'); + } + if (isdigit(c)) + { + index = (int)(c-'0'); + } + arma::mat data = arma::zeros(1, numVisible); + data[index] = 1; + + batch.insert_rows(batch.n_rows, data); + } + fclose(pFile); + + return batch; +} + int main() { - printf("Hallo, Welt!\n"); - - const string project("context99"); + const string project("poet"); Stack stack(".", project); -#if CREATE_TEST - stack.addTraining(stack.trainingBatch().row(1)); - printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); - - stack.delTraining(0); - printf("There are %d training samples\n", (int)stack.trainingBatch().n_rows); - - const int numLayers = 4; - int i = 0; - Layer *lowerLayer = new Layer("Layer", i, 16, 16, 8); - stack.addLayer(lowerLayer); - - for (++i; i < numLayers; i++) - { - Layer *layer = new Layer("Layer", i, lowerLayer->bh().n_elem, 1, lowerLayer->bh().n_elem >> 1); - layer->params().learningRate = lowerLayer->params().learningRate/2; - layer->params().numEpochs = lowerLayer->params().numEpochs/2; - lowerLayer = layer; - stack.addLayer(layer); - } - - // Save project - stack.save(); - - // Shake weights - stack.weightsInit(0.01); - - // Save weights - stack.saveWeights(); - -#else - // Load project - stack.load(); - - // Load weights - stack.loadWeights(); - - // Load training - stack.loadTrainingBatch(); - -#endif - -#if TRAIN_TEST - RbmListener statusDisplay; - - // Train stack - stack.train(&statusDisplay); - - // Save weights - stack.saveWeights(); -#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); - - v.print("v1"); - r = layer->upDownPass(v); - r.print("v2"); - - arma::mat err = Rbm::rms_error(r-v); - err.print("err"); - + arma::mat batch = createTraining("moby_ch1.txt"); + batch.save("poet.training.dat", arma::arma_ascii); return 0; }