From 9a8ad70a0be53f05498ec752493101283534f80b Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Tue, 18 Jan 2022 14:27:39 +0000 Subject: [PATCH] - improved RnnStack - AStack: context doesn't belong o training data git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@827 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/AStack.cpp | 15 ------------ source/RnnStack.cpp | 59 +++++++++++++++++++++++++++++++++------------ source/RnnStack.hpp | 15 +++++++++++- source/poet.cpp | 48 +++++++++++++++++++++++++++++++----- 4 files changed, 99 insertions(+), 38 deletions(-) diff --git a/source/AStack.cpp b/source/AStack.cpp index bfdbed4..c3c3978 100644 --- a/source/AStack.cpp +++ b/source/AStack.cpp @@ -177,21 +177,6 @@ size_t AStack::loadTrainingBatch(const std::string &dir, bool doNormalize) std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; } - // Migrate context part to training data - size_t numTraining = m_trainingBatch.n_rows; - size_t numVisible = getLayer(0)->numVisible(); - - if (m_trainingBatch.n_cols < numVisible) - { - size_t diff = numVisible - m_trainingBatch.n_cols; - arma::mat training_with_ctx = arma::join_rows(m_trainingBatch, arma::zeros(numTraining, diff)); - m_trainingBatch = training_with_ctx; - } - else if (m_trainingBatch.n_cols > numVisible) - { - m_trainingBatch = m_trainingBatch.submat(0, 0, numTraining-1, numVisible-1); - } - return m_trainingBatch.n_rows; } diff --git a/source/RnnStack.cpp b/source/RnnStack.cpp index c470cab..022c6ce 100644 --- a/source/RnnStack.cpp +++ b/source/RnnStack.cpp @@ -36,7 +36,7 @@ size_t RnnStack::numContext() arma::mat RnnStack::v_to_vc(const arma::mat& v) const { arma::mat c = arma::zeros(v.n_rows, m_numContext); - return arma::join_rows(v, c); + return v_to_vc(v, c); } arma::mat RnnStack::v_to_vc(const arma::mat& v, const arma::mat& c) const @@ -63,16 +63,18 @@ arma::mat RnnStack::vc_to_v(const arma::mat& vc) const arma::mat RnnStack::trainingBatchFrom(size_t layerId, const arma::mat& batch) { - arma::mat vc = batch; Layer *pLayer = getLayer(0); + arma::mat c = arma::zeros(batch.n_rows, m_numContext); + arma::mat v = vc_to_v(batch); + arma::mat vc = v_to_vc(v, c); while (pLayer) { if (layerId == pLayer->id()) { break; } - arma::mat c = pLayer->to_h_gibbs(vc); - arma::mat v = arma::shift(vc_to_v(vc), layerId, 1); + c = pLayer->to_h_gibbs(vc); + v = arma::shift(vc_to_v(vc), layerId, 1); vc = v_to_vc(v, c); pLayer = pLayer->next; } @@ -81,24 +83,49 @@ arma::mat RnnStack::trainingBatchFrom(size_t layerId, const arma::mat& batch) void RnnStack::train(const arma::mat& batch, Rbm::IListener* pListener) { - arma::mat thisBatch = batch; - Layer *pLayer = getLayer(0); - while (pLayer) + // thisBatch = {padding | batch} + int numTraining = batch.n_rows; + arma::mat padding = arma::zeros(getSeqLen()-1, batch.n_cols); + arma::mat batch_padded = arma::join_cols(padding, batch); + arma::mat c = arma::zeros(numTraining, m_numContext); + + for (int i=0; i < getSeqLen(); i++) { - thisBatch = trainingBatchFrom(pLayer->id(), thisBatch); - pLayer->train(thisBatch, pListener); + Layer *pLayer = getLayer(i); + int k = getSeqLen()-i-1; + arma::mat batch_shifted = batch_padded.rows(k, numTraining+k-1); + arma::mat vc = v_to_vc(batch_shifted, c); + pLayer->train(vc, pListener); + c = pLayer->to_h_gibbs(vc); pLayer = pLayer->next; } } void RnnStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) { - arma::mat thisBatch = batch; - Layer *pLayer = getLayer(layerId); - if (pLayer) - { - thisBatch = trainingBatchFrom(pLayer->id(), batch); - pLayer->train(thisBatch, pListener); - } } +arma::mat RnnStack::step_forward(arma::mat &state, const arma::mat& v_curr) +{ + arma::mat c = arma::zeros(1, m_numContext); + arma::mat z = arma::zeros(arma::size(v_curr)); + arma::mat r; + + if (state.is_empty()) + { + state = arma::zeros(getSeqLen(), v_curr.n_cols); + } + state = arma::shift(state, 1, 0); + state.row(0) = v_curr; + + for (int j=0; j < getSeqLen(); j++) + { + Layer *pLayer = getLayer(j); + arma::mat v = arma::join_rows(state.row(j), z); + arma::mat vc = v_to_vc(v, c); + c = pLayer->to_h_gibbs(vc); + vc = pLayer->to_v_gibbs(c); + r = to_next(vc_to_v(vc)); + } + return r; +} diff --git a/source/RnnStack.hpp b/source/RnnStack.hpp index 2d10ee8..094dbed 100644 --- a/source/RnnStack.hpp +++ b/source/RnnStack.hpp @@ -18,6 +18,8 @@ class RnnStack : public AStack { + const size_t NUM_CODES = 37; + public: RnnStack(const std::string &name, size_t numContext); RnnStack(const RnnStack& orig) = delete; @@ -34,9 +36,20 @@ public: arma::mat vc_to_v(const arma::mat &vc) const; arma::mat vc_to_c(const arma::mat &vc) const; + arma::mat step_forward(arma::mat &state, const arma::mat &v); + + arma::mat to_next(const arma::mat &v) + { + return v.cols(0, NUM_CODES-1); + } + + arma::mat to_curr(const arma::mat &v) + { + return v.cols(NUM_CODES, 2*NUM_CODES-1); + } + private: size_t m_numContext; - }; #endif /* RNNSTACK_HPP */ diff --git a/source/poet.cpp b/source/poet.cpp index 285fe88..ccc818d 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -35,7 +35,7 @@ class RbmListener : public Rbm::IListener const char punctuation[] = {' ', '.', '!', '?', 0}; const int NUM_CODES = 1 + 26 + 10; -const int SEQ_LENGTH = 5; +const int SEQ_LENGTH = 2; int char_is(char c, const char *pTable) { @@ -192,12 +192,14 @@ arma::mat to_next(arma::mat v) } #define CREATE_TRAINING 0 -#define DO_TRAINING 1 +#define DO_TRAINING 0 +#define DO_FORWARD 1 + int main() { #if CREATE_TRAINING arma::mat batch = createTraining("moby_ch1.txt", SEQ_LENGTH); - batch.save("poet.training.dat", arma::arma_ascii); + batch.save("poet2.training.dat", arma::arma_ascii); return 0; #endif @@ -210,19 +212,53 @@ int main() // Load training stack->loadTrainingBatch("."); + arma::mat t_vc = stack->trainingBatch(); + #if DO_TRAINING RbmListener listener; - arma::mat t_vc = stack->trainingBatch(); for (int i=0; i < stack->getSeqLen(); i++) { - stack->getLayer(i)->weightsInit(0.1,0); - stack->train(i, t_vc, &listener); +// stack->getLayer(i)->weightsInit(0.1,0); } + stack->train(t_vc, &listener); stack->saveWeights("."); stack->save("."); + printf("Curr\n"); + for (int i=0; i < t_vc.n_rows; i++) + { + arma::mat curr = stack->to_curr(t_vc.row(i)); + char c = idx2ch(curr.index_max()); + putchar(c); + } + printf("\n"); + printf("Next\n"); + for (int i=0; i < t_vc.n_rows; i++) + { + arma::mat next = stack->to_next(t_vc.row(i)); + char c = idx2ch(next.index_max()); + putchar(c); + } + printf("\n"); + +#endif + +#if DO_FORWARD + arma::mat state; + + for (int i=0; i < t_vc.n_rows; i++) + { + arma::mat curr = stack->to_curr(t_vc.row(i)); + arma::mat next = stack->to_next(t_vc.row(i)); + arma::mat r = stack->step_forward(state, curr); + char c = idx2ch(r.index_max()); + putchar(c); + } + printf("\n"); + #else + return 0; Layer *layer = stack->getLayer(0); int numTraining = stack->trainingBatch().n_rows;