- 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:
+4
-1
@@ -15,7 +15,10 @@
|
|||||||
#include "StackCreator.hpp"
|
#include "StackCreator.hpp"
|
||||||
#include <cassert>
|
#include <cassert>
|
||||||
|
|
||||||
|
#include "matutils.hpp"
|
||||||
|
|
||||||
using namespace std;
|
using namespace std;
|
||||||
|
using namespace Matutils;
|
||||||
|
|
||||||
const char *AStack::stackTypeStrings[NUM_STACKTYPES] = {"None", "Deep", "Rnn"};
|
const char *AStack::stackTypeStrings[NUM_STACKTYPES] = {"None", "Deep", "Rnn"};
|
||||||
|
|
||||||
@@ -158,7 +161,7 @@ size_t AStack::loadTrainingBatch(const std::string &dir, bool doNormalize)
|
|||||||
{
|
{
|
||||||
if (doNormalize)
|
if (doNormalize)
|
||||||
{
|
{
|
||||||
m_trainingBatch = Rbm::normalize(m_trainingBatch);
|
m_trainingBatch = normalize(m_trainingBatch);
|
||||||
}
|
}
|
||||||
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
|
std::cout << "Loaded " << m_trainingBatch.n_rows << " training samples\n";
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-49
@@ -13,8 +13,9 @@
|
|||||||
|
|
||||||
#include <cassert>
|
#include <cassert>
|
||||||
#include "Rbm.hpp"
|
#include "Rbm.hpp"
|
||||||
|
#include "matutils.hpp"
|
||||||
|
|
||||||
#define RBM_TRAIN_FLAT 0
|
using namespace Matutils;
|
||||||
|
|
||||||
Rbm::Rbm(size_t numVisible, size_t numHidden)
|
Rbm::Rbm(size_t numVisible, size_t numHidden)
|
||||||
: m_params()
|
: m_params()
|
||||||
@@ -272,27 +273,6 @@ arma::mat Rbm::prob(const arma::mat &src)
|
|||||||
return 1 / (1 + (arma::exp(-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
|
arma::mat Rbm::v_to_h(const arma::mat &visible) const
|
||||||
{
|
{
|
||||||
return visible * m_whv + arma::repmat(m_bhv, visible.n_rows, 1);
|
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);
|
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
|
const arma::mat& Rbm::whv() const
|
||||||
{
|
{
|
||||||
return m_whv;
|
return m_whv;
|
||||||
|
|||||||
+2
-7
@@ -124,7 +124,6 @@ public:
|
|||||||
|
|
||||||
void train(arma::mat const &batch, IListener *pListener=nullptr);
|
void train(arma::mat const &batch, IListener *pListener=nullptr);
|
||||||
|
|
||||||
static arma::mat normalize(const arma::mat &hidden);
|
|
||||||
const arma::mat& whv() const;
|
const arma::mat& whv() const;
|
||||||
const arma::mat& bv() const;
|
const arma::mat& bv() const;
|
||||||
const arma::mat& bh() const;
|
const arma::mat& bh() const;
|
||||||
@@ -148,18 +147,14 @@ public:
|
|||||||
static double rms_error_accu(arma::mat diffErr);
|
static double rms_error_accu(arma::mat diffErr);
|
||||||
static arma::mat rms_error(arma::mat diffErr);
|
static arma::mat rms_error(arma::mat diffErr);
|
||||||
|
|
||||||
private:
|
|
||||||
Params m_params;
|
|
||||||
arma::mat sample(arma::mat const &src);
|
|
||||||
void contrastiveDivergence(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv);
|
|
||||||
void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5);
|
|
||||||
|
|
||||||
protected:
|
protected:
|
||||||
|
Params m_params;
|
||||||
arma::mat m_bhv;
|
arma::mat m_bhv;
|
||||||
arma::mat m_bv;
|
arma::mat m_bv;
|
||||||
|
|
||||||
private:
|
private:
|
||||||
arma::mat m_whv;
|
arma::mat m_whv;
|
||||||
|
void contrastiveDivergence(arma::mat const &v_states, arma::mat &dwhv, arma::mat &dbhv, arma::mat &dbv);
|
||||||
|
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
/*
|
||||||
|
* To change this license header, choose License Headers in Project Properties.
|
||||||
|
* To change this template file, choose Tools | Templates
|
||||||
|
* and open the template in the editor.
|
||||||
|
*/
|
||||||
|
|
||||||
|
/*
|
||||||
|
* File: matutils.hpp
|
||||||
|
* Author: jens
|
||||||
|
*
|
||||||
|
* Created on 21. Januar 2022, 08:28
|
||||||
|
*/
|
||||||
|
|
||||||
|
#ifndef MATUTILS_HPP
|
||||||
|
#define MATUTILS_HPP
|
||||||
|
|
||||||
|
#include <math.h>
|
||||||
|
|
||||||
|
namespace Matutils
|
||||||
|
{
|
||||||
|
inline void uniform(arma::mat& srcDst, double stdDev=1.0, double mu=0.5)
|
||||||
|
{
|
||||||
|
#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
|
||||||
|
}
|
||||||
|
|
||||||
|
inline arma::mat 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
|
||||||
|
}
|
||||||
|
|
||||||
|
inline arma::mat 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;
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
#endif /* MATUTILS_HPP */
|
||||||
|
|
||||||
Reference in New Issue
Block a user