Files
jens 756c4d8670 - Noise_Init() depoends only on seed
- more detrmistic calls of shuffle()
- Shuffle() init weights with uniform noise

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@599 b431acfa-c32f-4a4a-93f1-934dc6c82436
2019-10-29 07:02:45 +00:00

289 lines
4.9 KiB
C++

/*
==============================================================================
Weights.hpp
Created: 21 Sep 2014 1:55:15pm
Author: jens
==============================================================================
*/
#ifndef WEIGHTS_HPP
#define WEIGHTS_HPP
#include <stdint.h>
#include <iostream>
#include <Eigen/Dense>
#include <stdio.h>
#include "noise.h"
using namespace std;
using namespace Eigen;
class Weights
{
public:
Weights(const char *pFilename)
: m_numVisible(0)
, m_numVisibleX(0)
, m_numVisibleY(0)
, m_numHidden(0)
{
Noise_Init(&m_noise, 0x32727155);
load(pFilename);
}
Weights(uint32_t numVisibleX, uint32_t numVisibleY, uint32_t numHidden)
: m_numVisible(0)
, m_numVisibleX(0)
, m_numVisibleY(0)
, m_numHidden(0)
{
Noise_Init(&m_noise, 0x32727155);
setUnits(numVisibleX, numVisibleY, numHidden);
}
Weights()
: m_numVisible(0)
, m_numVisibleX(0)
, m_numVisibleY(0)
, m_numHidden(0)
{
Noise_Init(&m_noise, 0x32727155);
}
Weights(const Weights &src)
: m_numVisible(0)
, m_numVisibleX(0)
, m_numVisibleY(0)
, m_numHidden(0)
{
Noise_Init(&m_noise, 0x32727155);
setUnits(src.m_numVisibleX, src.m_numVisibleY, src.m_numHidden);
*this = src;
}
~Weights()
{
Noise_Free(&m_noise);
free();
}
void setUnits(uint32_t numVisibleX, uint32_t numVisibleY, uint32_t numHidden)
{
if ((m_numVisibleX == numVisibleX) && (m_numVisibleY == numVisibleY) && (m_numHidden == numHidden))
{
return;
}
m_numVisibleX = numVisibleX;
m_numVisibleY = numVisibleY;
m_numVisible = numVisibleX * numVisibleY;
m_numHidden = numHidden;
m_w.resize(m_numVisible, m_numHidden);
m_bv.resize(m_numVisible);
m_bh.resize(m_numHidden);
}
void shuffle(double stdDev)
{
uint32_t i, j;
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
m_w(i,j) = stdDev*(Noise_Uniform(&m_noise) - 0.5);
}
}
for (i=0; i < m_numVisible; i++)
{
m_bv(i) = stdDev*(Noise_Uniform(&m_noise) - 0.5);
}
for (i=0; i < m_numHidden; i++)
{
m_bh(i) = stdDev*(Noise_Uniform(&m_noise) - 0.5);
}
}
Weights& operator= (const Weights &rhs)
{
m_bv = rhs.m_bv;
m_bh = rhs.m_bh;
m_w = rhs.m_w;
return *this;
}
MatrixXd& weights()
{
return m_w;
}
RowVectorXd& visibleBias()
{
return m_bv;
}
RowVectorXd& hiddenBias()
{
return m_bh;
}
void print()
{
uint32_t i, j;
printf("\n");
printf("w(v,h) = (v^, h>)\n");
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
printf("%3.6f ", m_w(i,j));
}
printf("\n");
}
printf("\n");
printf("bv = \n");
for (i=0; i < m_numVisible; i++)
{
printf("%3.6f\n", m_bv(i));
}
printf("\n");
printf("bh = \n");
for (i=0; i < m_numHidden; i++)
{
printf("%3.6f\n", m_bh(i));
}
printf("\n");
cout << m_w << endl;
}
uint32_t getNumVisible()
{
return m_numVisible;
}
uint32_t getNumVisibleX()
{
return m_numVisibleX;
}
uint32_t getNumVisibleY()
{
return m_numVisibleY;
}
uint32_t getNumHidden()
{
return m_numHidden;
}
void save(const char *pFilename)
{
FILE *pFile;
pFile = fopen(pFilename,"w");
if (!pFile)
return;
fprintf(pFile, "%d %d %d\n", m_numVisibleX, m_numVisibleY, m_numHidden);
uint32_t i, j;
for (i=0; i < m_numVisible; i++)
{
fprintf(pFile, "%3.6f\n", m_bv(i));
}
for (i=0; i < m_numHidden; i++)
{
fprintf(pFile, "%3.6f\n", m_bh(i));
}
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
fprintf(pFile, "%3.6f ", m_w(i,j));
}
fprintf(pFile, "\n");
}
fclose(pFile);
}
void load(const char *pFilename)
{
uint32_t numVisibleX;
uint32_t numVisibleY;
uint32_t numHidden;
FILE *pFile;
pFile = fopen(pFilename,"r");
if (!pFile)
return;
int result = fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden);
if (result < 0)
{
return;
}
setUnits(numVisibleX, numVisibleY, numHidden);
uint32_t i, j;
float v;
for (i=0; i < m_numVisible; i++)
{
result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_bv(i) = v;
}
}
for (i=0; i < m_numHidden; i++)
{
result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_bh(i) = v;
}
}
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
result = fscanf(pFile, "%f", &v);
if (result > 0)
{
m_w(i, j) = v;
}
}
}
fclose(pFile);
}
private:
uint32_t m_numVisible;
uint32_t m_numVisibleX;
uint32_t m_numVisibleY;
uint32_t m_numHidden;
noise_gen_t m_noise;
MatrixXd m_w;
RowVectorXd m_bv;
RowVectorXd m_bh;
void free()
{
setUnits(0, 0, 0);
}
};
#endif // WEIGHTS_HPP