From 11ac4650d960117070002e62c76830e334dd7639 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 14 Jan 2022 10:47:22 +0000 Subject: [PATCH] - more general to support more seq_lengths git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@807 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/poet.cpp | 104 ++++++++++++++++++++++++++++++++++++------------ 1 file changed, 78 insertions(+), 26 deletions(-) diff --git a/source/poet.cpp b/source/poet.cpp index 61fad72..e767009 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -33,7 +33,8 @@ class RbmListener : public Rbm::IListener }; const char punctuation[] = {' ', '.', '!', '?', 0}; -const int numCodes = 1 + 26 + 10; +const int NUM_CODES = 1 + 26 + 10; +const int SEQ_LENGTH = 5; int char_is(char c, const char *pTable) { @@ -84,7 +85,7 @@ char idx2ch(int index) return c; } -arma::mat createTraining(const string &filename) +arma::mat createTraining(const string &filename, size_t seq_len) { FILE *pFile; @@ -97,9 +98,11 @@ arma::mat createTraining(const string &filename) } int numTraining = 0; - int numVisible = 2*numCodes; + int numVisible = seq_len*NUM_CODES; arma::mat batch = arma::zeros(numTraining, numVisible); + arma::mat pattern = arma::zeros(seq_len, NUM_CODES); + int last_index = ch2idx(' '); while(!feof(pFile)) { @@ -115,37 +118,81 @@ arma::mat createTraining(const string &filename) { continue; } - arma::mat data = arma::zeros(1, numVisible); - - data[index+numCodes] = 1; - data[last_index] = 1; last_index = index; - batch.insert_rows(batch.n_rows, data); + pattern = arma::shift(pattern, 1); + arma::mat data = arma::zeros(1, NUM_CODES); + data[index] = 1; + pattern.row(0) = data; + arma::mat pattern_vector = pattern.as_row(); + batch.insert_rows(batch.n_rows, pattern_vector); } fclose(pFile); return batch; } +struct Rnn +{ + Rnn(Layer *layer) + : nV(layer->numVisible() - layer->context().n_cols) + , nH(layer->numHidden()) + , nVx(layer->numVisibleX()) + , nVy(layer->numVisibleY()) + , nC(layer->context().n_cols) + { + + } + size_t nV; + size_t nH; + size_t nVx; + size_t nVy; + size_t nC; + + arma::mat vcVec_to_vMat(arma::mat const &vcVec) + { + arma::mat vVec = vcVec.submat(0, 0, 0, nV-1); + arma::mat vMat = arma::reshape(vVec, nVx, nVy); + return vMat; + } + + arma::mat curr_vec(arma::mat const &vcVec) + { + arma::mat vMat = vcVec_to_vMat(vcVec); + + return vMat.col(nVy-1); + } + + arma::mat vcVec_next_step(arma::mat const &vcVec, arma::mat const &h) + { + arma::mat vMat = arma::shift(vcVec_to_vMat(vcVec), 1, 1); + arma::mat vVec = vMat.as_row(); + + arma::mat vcVec_next = arma::join_rows(vVec, h); + return vcVec_next; + } + +}; + arma::mat to_curr(arma::mat v) { - return v.cols(0, numCodes-1); + return v.cols(0, NUM_CODES-1); } arma::mat to_next(arma::mat v) { - return v.cols(numCodes, 2*numCodes-1); + return v.cols(NUM_CODES, 2*NUM_CODES-1); } #define CREATE_TRAINING 0 int main() { #if CREATE_TRAINING - arma::mat batch = createTraining("moby_ch1.txt"); + arma::mat batch = createTraining("moby_ch1.txt", SEQ_LENGTH); batch.save("poet.training.dat", arma::arma_ascii); + return 0; #endif - Stack stack(".", "poet2"); + Stack stack(".", "poet5"); // Load project stack.load(); @@ -159,45 +206,50 @@ int main() Layer *layer = stack.getLayer(0); int numTraining = stack.trainingBatch().n_rows; - arma::mat v = stack.trainingBatch(); + arma::mat t = stack.trainingBatch(); - arma::mat h = layer->toHiddenProbs(v); + arma::mat h = layer->toHiddenProbs(t); arma::mat r = layer->toVisibleProbs(h); layer->params().gibbsDoSampleHidden = false; layer->params().gibbsDoSampleVisible = false; + Rnn rnn(layer); + printf("\nStimulus: Training\n"); for (int i=0; i < numTraining; i++) { - arma::mat curr = to_curr(v.row(i)); + arma::mat curr = rnn.curr_vec(t.row(i)); char c = idx2ch(curr.index_max()); putchar(c); } + printf("\n"); - h = layer->toHiddenProbs(v.row(0)); - printf("\nStimulus: Next char and current h\n"); + h = layer->toHiddenProbs(t.row(0)); + printf("Stimulus: Next char and current h\n"); for (int i=1; i < numTraining+1; i++) { + arma::mat r; layer->gibbs_hv(h, r); - arma::mat curr = to_curr(r); - arma::mat next = to_next(r); - v = arma::join_rows(next, arma::zeros(1, next.n_cols), h); + arma::mat curr = rnn.curr_vec(r); + arma::mat v = rnn.vcVec_next_step(r, h); layer->gibbs_vh(v, h); char c = idx2ch(curr.index_max()); putchar(c); } - - h = layer->toHiddenProbs(v.row(0)); - printf("\nStimulus: Current h\n"); + printf("\n"); + + h = layer->toHiddenProbs(t.row(0)); + printf("Stimulus: Current h\n"); for (int i=1; i < numTraining+1; i++) { + arma::mat r; layer->gibbs_hv(h, r); - arma::mat curr = to_curr(r); - arma::mat next = to_next(r); - v = arma::join_rows(arma::zeros(1, next.n_cols), arma::zeros(1, next.n_cols), h); + arma::mat curr = rnn.curr_vec(r); + arma::mat v = rnn.vcVec_next_step(arma::zeros(1, r.n_cols), h); + layer->gibbs_vh(v, h); char c = idx2ch(curr.index_max());