diff --git a/Source/Rbm.cpp b/Source/Rbm.cpp index f74890e..1e7a801 100644 --- a/Source/Rbm.cpp +++ b/Source/Rbm.cpp @@ -66,7 +66,7 @@ MatrixXd Rbm::sample(MatrixXd const &src) noiseUniform(n); - return (src.array() > n.array()).cast(); + return (src.array() >= n.array()).cast(); } void Rbm::sample(MatrixXd &dst, MatrixXd const &src) @@ -86,6 +86,7 @@ RowVectorXd Rbm::probsLogistic(RowVectorXd const &src) MatrixXd Rbm::normalizeData(MatrixXd const &src) { +#if 0 double mean = src.array().mean(); cout << "mean" << ": " << endl << mean << endl; @@ -94,7 +95,19 @@ MatrixXd Rbm::normalizeData(MatrixXd const &src) double stddev = sqrt(x2.array().mean()); cout << "stddev" << ": " << endl << stddev << endl; - return x; + return x/stddev; + +#else + MatrixXd mean = src.colwise().mean(); + MatrixXd x = src - mean.replicate(src.rows(), 1); + MatrixXd x2 = x.cwiseProduct(x); + + + MatrixXd stddev_inv = x2.colwise().mean(); + stddev_inv = stddev_inv.array().sqrt().cwiseInverse(); + return x*stddev_inv.replicate(src.rows(), 1); + +#endif }