From 2a2be76b2c2b139d314f3babf8d80ead5c012368 Mon Sep 17 00:00:00 2001 From: jens Date: Wed, 31 Jan 2024 12:20:15 +0100 Subject: [PATCH] - AStack:Load show mean and stddev --- source/AStack.cpp | 3 +++ source/Rbm.cpp | 2 +- source/RbmComponent.cpp | 2 +- source/RnnComponentLayer.cpp | 2 +- source/matutils.hpp | 20 ++++++++++++++------ 5 files changed, 20 insertions(+), 9 deletions(-) diff --git a/source/AStack.cpp b/source/AStack.cpp index e824534..18daf2b 100644 --- a/source/AStack.cpp +++ b/source/AStack.cpp @@ -179,6 +179,9 @@ size_t AStack::loadTrainingBatch(const std::string &dir, bool doNormalize) { m_trainingBatch = normalize(m_trainingBatch); } + std::cout << "mean" << " : " << std::endl << mean(m_trainingBatch) << std::endl; + std::cout << "stddev" << ": " << std::endl << stddev(m_trainingBatch) << std::endl; + std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n"; } diff --git a/source/Rbm.cpp b/source/Rbm.cpp index ad33344..83e9bf2 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -146,7 +146,7 @@ void Rbm::cd(arma::mat const &v_data, arma::mat &dw, arma::mat &dbh, arma::mat & } else { - cd_hinton(v_data, dw, dbh, dbv); + cd_jens(v_data, dw, dbh, dbv); } } diff --git a/source/RbmComponent.cpp b/source/RbmComponent.cpp index b6ae2ca..e5fb615 100644 --- a/source/RbmComponent.cpp +++ b/source/RbmComponent.cpp @@ -318,7 +318,7 @@ void RbmComponent::onDownPass(const arma::mat& h) const DrawHidden->getData() = h; DrawHidden->DrawData(); - arma::mat r = Matutils::prob(h_to_v(h)); + arma::mat r = toVisibleProbs(h); reconstRedraw(r); } diff --git a/source/RnnComponentLayer.cpp b/source/RnnComponentLayer.cpp index 22bd1a3..beeb926 100644 --- a/source/RnnComponentLayer.cpp +++ b/source/RnnComponentLayer.cpp @@ -302,7 +302,7 @@ void RnnComponentLayer::onDownPass(const arma::mat& h) const DrawHidden->getData() = h; DrawHidden->DrawData(); - arma::mat r = Matutils::prob(h_to_v(h)); + arma::mat r = toVisibleProbs(h); reconstRedraw(r); } diff --git a/source/matutils.hpp b/source/matutils.hpp index a586af8..c6d5596 100644 --- a/source/matutils.hpp +++ b/source/matutils.hpp @@ -20,6 +20,7 @@ namespace Matutils { + const size_t NORMALIZING_DIM = 1; inline void uniform(arma::mat& srcDst, double stdDev=1.0, double mu=0.5) { #if 0 @@ -68,14 +69,13 @@ namespace Matutils // Dim = 0: Normalize over training all training pattern // Dim = 1: Normalize over single training pattern double k = 1; - size_t dim = 1; - arma::mat mean = arma::mean(src, dim); - arma::mat stddev = arma::stddev(src, 0, dim); + arma::mat mean = arma::mean(src, NORMALIZING_DIM); + arma::mat stddev = arma::stddev(src, 0, NORMALIZING_DIM); arma::mat mean_mat; arma::mat std_mat; - if (dim==0) + if (NORMALIZING_DIM==0) { mean_mat = arma::repmat(mean, src.n_rows, 1); std_mat = arma::repmat(stddev, src.n_rows, 1); @@ -88,11 +88,19 @@ namespace Matutils arma::mat xn = src - mean_mat; arma::mat y = xn/(std_mat + 1e-9); - std::cout << "mean" << " : " << std::endl << arma::mean(y, dim) << std::endl; - std::cout << "stddev" << ": " << std::endl << arma::stddev(y, 0, dim) << std::endl; return y; } + inline arma::mat mean(const arma::mat& src) + { + return arma::mean(src, NORMALIZING_DIM); + } + + inline arma::mat stddev(const arma::mat& src) + { + return arma::stddev(src, 0, NORMALIZING_DIM); + } + inline arma::mat char2vec(char c, size_t len) { arma::mat result = arma::zeros(1, len);