diff --git a/source/Layer.hpp b/source/Layer.hpp index b3ced1b..530137b 100644 --- a/source/Layer.hpp +++ b/source/Layer.hpp @@ -76,9 +76,54 @@ public: return m_numVisibleY; } - int numHidden() + arma::mat gibbsPass(arma::mat &vr) { - return whv().n_cols; + arma::mat h; + for (int i=0; i < params().numGibbs; i++) + { + h = prob(v_to_h(vr)); + vr = prob(h_to_v(h)); + } + return h; + } + + arma::mat upPass(arma::mat const &v) + { + arma::mat r = v; + arma::mat h = gibbsPass(r); + if (next) + { + next->upPass(h); + } + else + { + return h; + } + } + + arma::mat downPass(arma::mat const &h) + { + arma::mat v = prob(h_to_v(h)); + if (prev) + { + return prev->downPass(v); + } + return v; + } + + arma::mat upDownPass(arma::mat const &v) + { + arma::mat r = v; + arma::mat h = gibbsPass(r); + if (next) + { + next->upDownPass(h); + } + else if (prev) + { + prev->downPass(r); + } + return r; } bool weightsLoad(std::string const &dir, std::string const &prj); diff --git a/source/RbmComponent.cpp b/source/RbmComponent.cpp index c434fba..be203e9 100644 --- a/source/RbmComponent.cpp +++ b/source/RbmComponent.cpp @@ -349,13 +349,9 @@ void RbmComponent::upDownPass(const arma::mat& vc) } else if (prev) { - if (prev) - { - RbmComponent *pComp = static_cast (prev); - pComp->downPass(getReconst()); - } + RbmComponent *pComp = static_cast (prev); + pComp->downPass(getReconst()); } - } arma::mat RbmComponent::getConvolutedWeight(arma::mat const &w) diff --git a/source/main.cpp b/source/main.cpp index 098dfd0..7319a77 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -84,13 +84,13 @@ int main() stack.loadTrainingBatch(); +#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); -#if 1 const int numLayers = 4; int i = 0; Layer *lowerLayer = new Layer("Layer", i, 16, 16, 8); @@ -132,5 +132,10 @@ int main() arma::mat v = arma::randu(stack.trainingBatch().n_rows, layer->bv().n_elem); arma::mat h = layer->toHiddenProbs(v); arma::mat r = layer->toVisibleProbs(h); + + v.print("v1"); + r = layer->upDownPass(v); + r.print("v2"); + return 0; }