diff --git a/source/AStack.cpp b/source/AStack.cpp index c3c3978..9c4de3d 100644 --- a/source/AStack.cpp +++ b/source/AStack.cpp @@ -141,22 +141,6 @@ bool AStack::saveWeights(const std::string &dir) return true; } -arma::mat AStack::trainingBatchFrom(size_t layerId, const arma::mat& batch) -{ - arma::mat thisBatch = batch; - Layer *pLayer = getLayer(0); - while (pLayer) - { - if (pLayer->id() == layerId) - { - break; - } - thisBatch = pLayer->toHiddenProbs(thisBatch); - pLayer = pLayer->next; - } - return thisBatch; -} - arma::mat& AStack::trainingBatch() { return m_trainingBatch; diff --git a/source/AStack.hpp b/source/AStack.hpp index d1d06e6..a09d01f 100644 --- a/source/AStack.hpp +++ b/source/AStack.hpp @@ -69,9 +69,7 @@ public: bool loadWeights(const std::string &dir); bool saveWeights(const std::string &dir); - virtual void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) = 0; virtual void train(const arma::mat& batch, Rbm::IListener* pListener) = 0; - virtual arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch) = 0; size_t numTraining(); void addTraining(const arma::mat &toAdd); diff --git a/source/DeepStack.hpp b/source/DeepStack.hpp index 9319c2c..0cc4b2f 100644 --- a/source/DeepStack.hpp +++ b/source/DeepStack.hpp @@ -28,9 +28,9 @@ public: DeepStack(const DeepStack& orig) = delete; virtual ~DeepStack(); - arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch); void train(const arma::mat& batch, Rbm::IListener* pListener) override; - void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override; + arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch); + void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener); arma::mat upPass(size_t layerId, arma::mat const &v); arma::mat downPass(size_t layerId, arma::mat const &h); diff --git a/source/RnnStack.cpp b/source/RnnStack.cpp index 022c6ce..e72f031 100644 --- a/source/RnnStack.cpp +++ b/source/RnnStack.cpp @@ -61,26 +61,6 @@ arma::mat RnnStack::vc_to_v(const arma::mat& vc) const return vc.submat(0, 0, vc.n_rows - 1, numVisible - m_numContext - 1); } -arma::mat RnnStack::trainingBatchFrom(size_t layerId, const arma::mat& 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; - } - c = pLayer->to_h_gibbs(vc); - v = arma::shift(vc_to_v(vc), layerId, 1); - vc = v_to_vc(v, c); - pLayer = pLayer->next; - } - return vc; -} - void RnnStack::train(const arma::mat& batch, Rbm::IListener* pListener) { // thisBatch = {padding | batch} @@ -101,10 +81,6 @@ void RnnStack::train(const arma::mat& batch, Rbm::IListener* pListener) } } -void RnnStack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) -{ -} - arma::mat RnnStack::step_forward(arma::mat &state, const arma::mat& v_curr) { arma::mat c = arma::zeros(1, m_numContext); @@ -121,11 +97,11 @@ arma::mat RnnStack::step_forward(arma::mat &state, const arma::mat& 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 v = arma::join_rows(z, state.row(j)); 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)); + pLayer->gibbs_vh(vc, c); + + r = vc_to_v(vc); } return r; } diff --git a/source/RnnStack.hpp b/source/RnnStack.hpp index 094dbed..2541cba 100644 --- a/source/RnnStack.hpp +++ b/source/RnnStack.hpp @@ -25,9 +25,7 @@ public: RnnStack(const RnnStack& orig) = delete; virtual ~RnnStack(); - arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch) override; void train(const arma::mat& batch, Rbm::IListener* pListener) override; - void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) override; size_t getSeqLen(); size_t numContext(); @@ -48,6 +46,17 @@ public: 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: size_t m_numContext; }; diff --git a/source/poet.cpp b/source/poet.cpp index ccc818d..104f543 100644 --- a/source/poet.cpp +++ b/source/poet.cpp @@ -137,60 +137,6 @@ arma::mat createTraining(const string &filename, size_t seq_len) return batch; } -struct Rnn -{ - Rnn(Layer *layer) - : nV(layer->numVisible() - layer->numHidden()) - , nH(layer->numHidden()) - , nVx(layer->numVisibleX()) - , nVy(layer->numVisibleY()) - , nC(layer->numHidden()) - { - - } - 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(0); - } - - arma::mat vcVec_next_step(arma::mat const &vCurr, arma::mat const &vcVec_last, arma::mat const &h) - { - arma::mat vMat_last = arma::shift(vcVec_to_vMat(vcVec_last), 1, 1); - vMat_last.col(0) = arma::zeros(nVx, 1); - vMat_last.col(1) = vCurr; - - arma::mat vVec = vMat_last.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, NUM_CODES-1); -} - -arma::mat to_next(arma::mat v) -{ - return v.cols(NUM_CODES, 2*NUM_CODES-1); -} - #define CREATE_TRAINING 0 #define DO_TRAINING 0 #define DO_FORWARD 1 @@ -211,17 +157,36 @@ int main() // Load training stack->loadTrainingBatch("."); - arma::mat t_vc = stack->trainingBatch(); + Rbm::Params params = stack->getLayer(0)->params(); + #if DO_TRAINING + int numGibbs = params.numGibbs; + params.numGibbs = 1; + stack->setParams(params); RbmListener listener; for (int i=0; i < stack->getSeqLen(); i++) { -// stack->getLayer(i)->weightsInit(0.1,0); + stack->getLayer(i)->weightsInit(0.01,0); } + + // Phase 10 + params.numEpochs = 2000; + stack->setParams(params); stack->train(t_vc, &listener); + + // Phase 11 + params.numEpochs = 10000; + stack->setParams(params); + stack->train(t_vc, &listener); + + // Phase 20 + params.numGibbs = numGibbs; + stack->setParams(params); + stack->train(t_vc, &listener); + stack->saveWeights("."); stack->save("."); @@ -244,81 +209,21 @@ int main() #endif -#if DO_FORWARD + params.gibbsDoSampleHidden = false; + params.gibbsDoSampleVisible = false; + stack->setParams(params); 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()); + arma::mat next = stack->to_next(r); + char c = idx2ch(next.index_max()); putchar(c); } printf("\n"); -#else - return 0; - Layer *layer = stack->getLayer(0); - int numTraining = stack->trainingBatch().n_rows; - - arma::mat t = stack->trainingBatch(); - - 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"); - h = layer->toHiddenProbs(t); - r = t; - for (int i=0; i < numTraining; i++) - { - arma::mat r = t.row(i); - layer->gibbs_vh(r, h); - layer->gibbs_hv(h, r); - arma::mat curr = rnn.curr_vec(r); - char c = idx2ch(curr.index_max()); - putchar(c); - } - printf("\n"); - - arma::mat v = t.row(0); - printf("Stimulus: Next char and current h\n"); - for (int i=1; i < numTraining+1; i++) - { - arma::mat r = v; - layer->gibbs_vh(r, h); - arma::mat curr = rnn.curr_vec(r); - v = rnn.vcVec_next_step(curr, r, h); - - layer->gibbs_vh(v, h); - - char c = idx2ch(curr.index_max()); - putchar(c); - } - printf("\n"); - - v = t.row(0); - h = layer->toHiddenProbs(v); - printf("Stimulus: Current h\n"); - for (int i=1; i < numTraining+1; i++) - { - arma::mat r = v; - layer->gibbs_hv(h, r); - arma::mat curr = rnn.curr_vec(r); - v = rnn.vcVec_next_step(curr, arma::zeros(1, r.n_cols), h); - - layer->gibbs_vh(v, h); - - char c = idx2ch(curr.index_max()); - putchar(c); - } - -#endif - printf("\n\nEnd of program\n"); + printf("\nEnd of program\n"); return 0; }