- refactored common functions into matutils

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@855 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2022-01-21 07:46:40 +00:00
parent c10d9c6e38
commit aef3049856
4 changed files with 80 additions and 57 deletions
+2 -49
View File
@@ -13,8 +13,9 @@
#include <cassert>
#include "Rbm.hpp"
#include "matutils.hpp"
#define RBM_TRAIN_FLAT 0
using namespace Matutils;
Rbm::Rbm(size_t numVisible, size_t numHidden)
: m_params()
@@ -272,27 +273,6 @@ arma::mat Rbm::prob(const arma::mat &src)
return 1 / (1 + (arma::exp(-src)));
}
arma::mat Rbm::sample(const arma::mat &src)
{
arma::mat dst = src;
uniform(dst);
#if 0
for (size_t i=0; i < src.n_rows; i++)
{
for (size_t j=0; j < src.n_cols; j++)
{
dst(i, j) = src(i, j) >= dst(i, j);
}
}
return dst;
#else
arma::umat res = (dst < src);
return arma::conv_to<arma::mat>::from(res);
#endif
}
arma::mat Rbm::v_to_h(const arma::mat &visible) const
{
return visible * m_whv + arma::repmat(m_bhv, visible.n_rows, 1);
@@ -303,33 +283,6 @@ arma::mat Rbm::h_to_v(const arma::mat &hidden) const
return hidden * m_whv.t() + arma::repmat(m_bv, hidden.n_rows, 1);
}
arma::mat Rbm::normalize(const arma::mat& src)
{
double mean = arma::accu(src)/src.n_elem;
arma::mat x = src - mean;
arma::mat x2 = x % x;
double stddev = sqrt(arma::accu(x2)/x2.n_elem);
std::cout << "mean" << " : " << std::endl << mean << std::endl;
std::cout << "stddev" << ": " << std::endl << stddev << std::endl;
return x/stddev;
}
void Rbm::uniform(arma::mat& srcDst, double stdDev, double mu)
{
#if 0
for (size_t i=0; i < srcDst.n_rows; i++)
{
for (size_t j=0; j < srcDst.n_cols; j++)
{
srcDst(i, j) = stdDev*(Noise_Uniform(&m_noise) + mu - 0.5);
}
}
#else
srcDst = stdDev*(arma::randu(arma::size(srcDst)) + mu - 0.5);
#endif
}
const arma::mat& Rbm::whv() const
{
return m_whv;