- use expectations

- improved gibbs sampling
- LayerArray is template class

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@18 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2014-10-06 17:08:50 +00:00
parent f62f1283f8
commit fc758b853a
122 changed files with 1053 additions and 1149 deletions
+55 -15
View File
@@ -10,7 +10,7 @@
#ifndef WEIGHTS_HPP
#define WEIGHTS_HPP
#include <cstdint>
#include <stdint.h>
#include "noise.h"
class Weights
@@ -33,9 +33,8 @@ public:
, m_numVisible(numVisible)
, m_numHidden(numHidden)
{
alloc();
Noise_Init(&m_noise, 0x32727155);
alloc(numVisible, numHidden);
shuffle(0);
}
@@ -49,6 +48,18 @@ public:
Noise_Init(&m_noise, 0x32727155);
}
Weights(const Weights &src)
: m_ppW(nullptr)
, m_pBiasVisible(nullptr)
, m_pBiasHidden(nullptr)
, m_numVisible(0)
, m_numHidden(0)
{
Noise_Init(&m_noise, 0x32727155);
alloc(src.m_numVisible, src.m_numHidden);
*this = src;
}
~Weights()
{
Noise_Free(&m_noise);
@@ -57,9 +68,7 @@ public:
void setUnits(uint32_t numVisible, uint32_t numHidden)
{
m_numVisible = numVisible;
m_numHidden = numHidden;
alloc();
alloc(numVisible, numHidden);
shuffle(0);
}
@@ -87,6 +96,30 @@ public:
}
}
Weights& operator= (const Weights &rhs)
{
uint32_t i, j;
for (j=0; j < m_numVisible; j++)
{
m_pBiasVisible[j] = rhs.m_pBiasVisible[j];
}
for (i=0; i < m_numHidden; i++)
{
m_pBiasHidden[i] = rhs.m_pBiasHidden[i];
}
for (i=0; i < m_numHidden; i++)
{
for (j=0; j < m_numVisible; j++)
{
m_ppW[i][j] = rhs.m_ppW[i][j];
}
}
return *this;
}
double **getWeights() const
{
return m_ppW;
@@ -183,6 +216,8 @@ public:
void load(const char *pFilename)
{
uint32_t numVisible;
uint32_t numHidden;
FILE *pFile;
pFile = fopen(pFilename,"r");
@@ -190,26 +225,25 @@ public:
if (!pFile)
return;
m_numVisible = m_numHidden = 0;
fscanf(pFile, "%d %d\n", &m_numVisible, &m_numHidden);
fscanf(pFile, "%d %d\n", &numVisible, &numHidden);
alloc();
alloc(numVisible, numHidden);
uint32_t i, j;
float v;
for (i=0; i < m_numVisible; i++)
for (i=0; i < numVisible; i++)
{
fscanf(pFile, "%f", &v);
m_pBiasVisible[i] = v;
}
for (i=0; i < m_numHidden; i++)
for (i=0; i < numHidden; i++)
{
fscanf(pFile, "%f", &v);
m_pBiasHidden[i] = v;
}
for (i=0; i < m_numVisible; i++)
for (i=0; i < numVisible; i++)
{
for (j=0; j < m_numHidden; j++)
for (j=0; j < numHidden; j++)
{
fscanf(pFile, "%f", &v);
@@ -228,17 +262,23 @@ private:
uint32_t m_numHidden;
noise_gen_t m_noise;
void alloc()
void alloc(uint32_t numVisible, uint32_t numHidden)
{
uint32_t i;
if (m_ppW)
{
if ((numVisible == m_numVisible) && (numHidden == m_numHidden))
{
return;
}
free();
alloc();
alloc(numVisible, numHidden);
}
else
{
m_numVisible = numVisible;
m_numHidden = numHidden;
m_ppW = new double*[m_numHidden];
for (i=0; i < m_numHidden; i++)
{