Files
Rbm-legacy/Source/Weights.hpp
T
jens f62f1283f8 further development
git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@17 b431acfa-c32f-4a4a-93f1-934dc6c82436
2014-10-04 19:09:35 +00:00

276 lines
4.4 KiB
C++

/*
==============================================================================
Weights.hpp
Created: 21 Sep 2014 1:55:15pm
Author: jens
==============================================================================
*/
#ifndef WEIGHTS_HPP
#define WEIGHTS_HPP
#include <cstdint>
#include "noise.h"
class Weights
{
public:
Weights(const char *pFilename)
: m_ppW(nullptr)
, m_pBiasVisible(nullptr)
, m_pBiasHidden(nullptr)
, m_numVisible(0)
, m_numHidden(0)
{
load(pFilename);
}
Weights(uint32_t numVisible, uint32_t numHidden)
: m_ppW(nullptr)
, m_pBiasVisible(nullptr)
, m_pBiasHidden(nullptr)
, m_numVisible(numVisible)
, m_numHidden(numHidden)
{
alloc();
Noise_Init(&m_noise, 0x32727155);
shuffle(0);
}
Weights()
: m_ppW(nullptr)
, m_pBiasVisible(nullptr)
, m_pBiasHidden(nullptr)
, m_numVisible(0)
, m_numHidden(0)
{
Noise_Init(&m_noise, 0x32727155);
}
~Weights()
{
Noise_Free(&m_noise);
free();
}
void setUnits(uint32_t numVisible, uint32_t numHidden)
{
m_numVisible = numVisible;
m_numHidden = numHidden;
alloc();
shuffle(0);
}
void shuffle(double stdDev)
{
uint32_t i, j;
double kdev = stdDev*sqrt(12.0);
for (j=0; j < m_numVisible; j++)
{
m_pBiasVisible[j] = kdev*Noise_Uniform(&m_noise, 0.5);
}
for (i=0; i < m_numHidden; i++)
{
m_pBiasHidden[i] = kdev*Noise_Uniform(&m_noise, 0.5);
}
for (i=0; i < m_numHidden; i++)
{
for (j=0; j < m_numVisible; j++)
{
m_ppW[i][j] = kdev*Noise_Uniform(&m_noise, 0.5);
}
}
}
double **getWeights() const
{
return m_ppW;
}
double *getBiasVisible() const
{
return m_pBiasVisible;
}
double *getBiasHidden() const
{
return m_pBiasHidden;
}
void print()
{
uint32_t i, j;
double w;
printf("\n");
printf("w(v,h) = (v^, h>)\n");
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
w = m_ppW[j][i];
printf("%3.6f ", w);
}
printf("\n");
}
printf("\n");
printf("bv = \n");
for (i=0; i < m_numVisible; i++)
{
w = m_pBiasVisible[i];
printf("%3.6f\n", w);
}
printf("\n");
printf("bh = \n");
for (i=0; i < m_numHidden; i++)
{
w = m_pBiasHidden[i];
printf("%3.6f\n", w);
}
printf("\n");
}
uint32_t getNumVisible()
{
return m_numVisible;
}
uint32_t getNumHidden()
{
return m_numHidden;
}
void save(const char *pFilename)
{
FILE *pFile;
pFile = fopen(pFilename,"w");
if (!pFile)
return;
fprintf(pFile, "%d %d\n", m_numVisible, m_numHidden);
uint32_t i, j;
for (i=0; i < m_numVisible; i++)
{
fprintf(pFile, "%3.6f\n", m_pBiasVisible[i]);
}
for (i=0; i < m_numHidden; i++)
{
fprintf(pFile, "%3.6f\n", m_pBiasHidden[i]);
}
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
fprintf(pFile, "%3.6f ", m_ppW[j][i]);
}
fprintf(pFile, "\n");
}
fclose(pFile);
}
void load(const char *pFilename)
{
FILE *pFile;
pFile = fopen(pFilename,"r");
if (!pFile)
return;
m_numVisible = m_numHidden = 0;
fscanf(pFile, "%d %d\n", &m_numVisible, &m_numHidden);
alloc();
uint32_t i, j;
float v;
for (i=0; i < m_numVisible; i++)
{
fscanf(pFile, "%f", &v);
m_pBiasVisible[i] = v;
}
for (i=0; i < m_numHidden; i++)
{
fscanf(pFile, "%f", &v);
m_pBiasHidden[i] = v;
}
for (i=0; i < m_numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
{
fscanf(pFile, "%f", &v);
m_ppW[j][i] = v;
}
}
fclose(pFile);
}
private:
double **m_ppW;
double *m_pBiasVisible;
double *m_pBiasHidden;
uint32_t m_numVisible;
uint32_t m_numHidden;
noise_gen_t m_noise;
void alloc()
{
uint32_t i;
if (m_ppW)
{
free();
alloc();
}
else
{
m_ppW = new double*[m_numHidden];
for (i=0; i < m_numHidden; i++)
{
m_ppW[i] = new double[m_numVisible];
}
m_pBiasVisible = new double[m_numVisible];
m_pBiasHidden = new double[m_numHidden];
}
}
void free()
{
uint32_t i;
if (m_ppW)
{
for (i=0; i < m_numHidden; i++)
{
delete [] m_ppW[i];
}
delete [] m_ppW;
m_ppW = nullptr;
}
delete [] m_pBiasVisible;
m_pBiasVisible = nullptr;
delete [] m_pBiasHidden;
m_pBiasHidden = nullptr;
}
};
#endif // WEIGHTS_HPP