From e58593973dce7cea08d0f1c347f5f1be7f84bfc2 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 21 Jan 2022 16:41:11 +0000 Subject: [PATCH] - refactored - create training one th e fly from file name git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@860 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/RnnStack.cpp | 48 +++++++++++++++++++++++++++++++++ source/RnnStack.hpp | 58 ++++++---------------------------------- source/RnnTextHelper.cpp | 16 +++++------ source/poet.cpp | 39 ++++++++++++--------------- 4 files changed, 81 insertions(+), 80 deletions(-) diff --git a/source/RnnStack.cpp b/source/RnnStack.cpp index d9de7e8..6c34cc5 100644 --- a/source/RnnStack.cpp +++ b/source/RnnStack.cpp @@ -84,3 +84,51 @@ arma::mat RnnStack::step_forward(arma::mat &state, const arma::mat& v_curr) } return r; } + +void RnnStack::clamp_one_hot(arma::mat& srcDst) +{ + int index = srcDst.index_max(); + srcDst = arma::zeros(arma::size(srcDst)); + srcDst[index] = 1; +} + +void RnnStack::sample_one_hot(arma::mat& srcDst) +{ + double k = arma::accu(srcDst); + + if (k != 0) + { + srcDst = srcDst / k; + } + + arma::mat ps = Matutils::sample(srcDst); + if (arma::accu(ps) > 1) + { + arma::uvec q1 = find(ps > 0); + int winner = (q1.n_elem - 1) * arma::randu(1)[0]; + int wix = q1(winner); + + srcDst = zeros(arma::size(srcDst)); + srcDst[wix] = 1; + } +} + +arma::mat RnnStack::to_curr(const arma::mat& v) +{ + return v.cols(NUM_CODES, 2 * NUM_CODES - 1); +} + +arma::mat RnnStack::to_next(const arma::mat& v) +{ + return v.cols(0, NUM_CODES - 1); +} + +void RnnStack::setParams(const Rbm::Params& param) +{ + Layer *pLayer = getLayer(0); + while (pLayer) + { + pLayer->params() = param; + pLayer = pLayer->next; + } +} diff --git a/source/RnnStack.hpp b/source/RnnStack.hpp index a9751df..f761bba 100644 --- a/source/RnnStack.hpp +++ b/source/RnnStack.hpp @@ -26,63 +26,21 @@ public: RnnStack(const RnnStack& orig) = delete; virtual ~RnnStack(); + size_t getSeqLen(); + void train(const arma::mat& batch, Rbm::IListener* pListener) override; - size_t getSeqLen(); arma::mat v_to_vc(const arma::mat &v, const arma::mat &c) const; arma::mat vc_to_v(const arma::mat &vc, size_t numContext) const; - arma::mat step_forward(arma::mat &state, const arma::mat &v); - void sample_one_hot(arma::mat &srcDst) - { - double k = arma::accu(srcDst); - - if (k != 0) - { - srcDst = srcDst / k; - } - - arma::mat ps = Matutils::sample(srcDst); - if (arma::accu(ps) > 1) - { - arma::uvec q1 = find(ps > 0); - int winner = (q1.n_elem-1) * arma::randu(1)[0]; - int wix = q1(winner); - - srcDst = zeros(arma::size(srcDst)); - srcDst[wix] = 1; - } - } - - void clamp_one_hot(arma::mat &srcDst) - { - int index = srcDst.index_max(); - srcDst = arma::zeros(arma::size(srcDst)); - srcDst[index] = 1; - } - - arma::mat to_next(const arma::mat &v) - { - return v.cols(0, NUM_CODES-1); - } + void sample_one_hot(arma::mat &srcDst); + void clamp_one_hot(arma::mat &srcDst); + arma::mat to_next(const arma::mat &v); + arma::mat to_curr(const arma::mat &v); + + void setParams(Rbm::Params const ¶m); - arma::mat to_curr(const arma::mat &v) - { - return v.cols(NUM_CODES, 2*NUM_CODES-1); - } - - void setParams(Rbm::Params const ¶m) - { - Layer *pLayer = getLayer(0); - while (pLayer) - { - pLayer->params() = param; - pLayer = pLayer->next; - } - } - - private: }; diff --git a/source/RnnTextHelper.cpp b/source/RnnTextHelper.cpp index 13310c5..4567a63 100644 --- a/source/RnnTextHelper.cpp +++ b/source/RnnTextHelper.cpp @@ -12,6 +12,7 @@ */ #include "RnnTextHelper.hpp" +#include "matutils.hpp" using namespace std; @@ -74,6 +75,7 @@ char RnnTextHelper::idx2ch(int index) arma::mat RnnTextHelper::createTraining(const string& filename, size_t seq_len) { + const char SPACE = ' '; FILE *pFile; pFile = fopen(filename.c_str(), "r"); @@ -92,33 +94,31 @@ arma::mat RnnTextHelper::createTraining(const string& filename, size_t seq_len) for (int i = 0; i < seq_len; i++) { - pattern.row(i)[0] = 1; + pattern.row(i) = Matutils::char2vec(SPACE, NUM_CODES); } - int last_index = ch2idx(' '); + int last_index = ch2idx(SPACE); while (!feof(pFile)) { - char c; - int result = fread(&c, 1, 1, pFile); + char ch = SPACE; + int result = fread(&ch, 1, 1, pFile); if (result < 0) { break; } - int index = ch2idx(c); + int index = ch2idx(ch); if (index == 0 and last_index == 0) { continue; } last_index = index; - arma::mat data = arma::zeros(1, NUM_CODES); - data[index] = 1; if (seq_len > 1) { pattern = arma::shift(pattern, 1); } - pattern.row(0) = data; + pattern.row(0) = Matutils::char2vec(ch, NUM_CODES); arma::mat pattern_vector = pattern.as_row(); batch.insert_rows(batch.n_rows, pattern_vector); } diff --git a/source/poet.cpp b/source/poet.cpp index 8dff32a..e83c0a4 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -42,21 +42,18 @@ class RbmListener : public Rbm::IListener std::cout << "L1 = " << status.L1 << std::endl; std::cout << "L2 = " << status.L2 << std::endl; m_last_progress = status.progress; - forward(DEFAULT_START_WORD); + forward(DEFAULT_START_WORD, 100); } - return true; } - void forward(std::string const &start) + void forward(std::string const &start, size_t len) { RnnStack *stack = reinterpret_cast(&m_stack); arma::mat state; arma::mat curr; arma::mat next; std::string curr_str; - std::string next_str; - arma::mat t_vc = stack->trainingBatch(); for (int i=0; i < start.size(); i++) { @@ -68,7 +65,7 @@ class RbmListener : public Rbm::IListener } cout << "Start: " << curr_str << std::endl; - for (int i=stack->getSeqLen(); i < t_vc.n_rows; i++) + for (int i=stack->getSeqLen(); i < len; i++) { curr = next; curr_str.append(1, Matutils::vec2char(curr)); @@ -110,23 +107,12 @@ int main(int argc, char *argv[]) } } - if (command == Command::Create) - { - arma::mat batch = RnnTextHelper::createTraining("moby_ch1.txt", RnnTextHelper::SEQ_LENGTH); - batch.save("poet2.training.dat", arma::arma_ascii); - return 0; - } - // Load project RnnStack *stack = reinterpret_cast(StackCreator::fromFile(".", "poet_2v_5s")); // Load weights stack->loadWeights("."); - // Load training - stack->loadTrainingBatch("."); - arma::mat t_vc = stack->trainingBatch(); - RbmListener listener(*stack); if (command == Command::Reset) @@ -137,22 +123,31 @@ int main(int argc, char *argv[]) if (command == Command::Train) { + arma::mat batch; + if (argv[2] != nullptr) + { + batch = RnnTextHelper::createTraining(argv[2], RnnTextHelper::SEQ_LENGTH); + } + else + { + batch = RnnTextHelper::createTraining("batch1.txt", RnnTextHelper::SEQ_LENGTH); + } + // Phase 10 - stack->train(t_vc, &listener); + stack->train(batch, &listener); stack->saveWeights("."); #if DO_SAVE_PROJECT_AFTER_TRAINING stack->save("."); #endif printf("Curr\n"); - for (int i=0; i < t_vc.n_rows; i++) + for (int i=0; i < batch.n_rows; i++) { - arma::mat curr = stack->to_curr(t_vc.row(i)); + arma::mat curr = stack->to_curr(batch.row(i)); char c = Matutils::vec2char(curr); putchar(c); } printf("\n"); - command = Command::Forward; } if (command == Command::Forward) @@ -163,7 +158,7 @@ int main(int argc, char *argv[]) { start = std::string(argv[2]); } - listener.forward(start); + listener.forward(start, 100); } printf("\nEnd of program\n"); return 0;