[RBM]
- committed local changes git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@270 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+78
-35
@@ -24,22 +24,28 @@ 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 numVisible, uint32_t numHidden)
|
||||
: m_numVisible(numVisible)
|
||||
, m_numHidden(numHidden)
|
||||
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);
|
||||
alloc(numVisible, numHidden);
|
||||
setUnits(numVisibleX, numVisibleY, numHidden);
|
||||
}
|
||||
|
||||
Weights()
|
||||
: m_numVisible(0)
|
||||
, m_numVisibleX(0)
|
||||
, m_numVisibleY(0)
|
||||
, m_numHidden(0)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
@@ -47,10 +53,12 @@ public:
|
||||
|
||||
Weights(const Weights &src)
|
||||
: m_numVisible(0)
|
||||
, m_numVisibleX(0)
|
||||
, m_numVisibleY(0)
|
||||
, m_numHidden(0)
|
||||
{
|
||||
Noise_Init(&m_noise, 0x32727155);
|
||||
alloc(src.m_numVisible, src.m_numHidden);
|
||||
setUnits(src.m_numVisibleX, src.m_numVisibleY, src.m_numHidden);
|
||||
*this = src;
|
||||
}
|
||||
|
||||
@@ -60,9 +68,23 @@ public:
|
||||
free();
|
||||
}
|
||||
|
||||
void setUnits(uint32_t numVisible, uint32_t numHidden)
|
||||
void setUnits(uint32_t numVisibleX, uint32_t numVisibleY, uint32_t numHidden)
|
||||
{
|
||||
alloc(numVisible, numHidden);
|
||||
shuffle(0);
|
||||
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_sigma.resize(m_numVisible);
|
||||
m_mean.resize(m_numVisible);
|
||||
m_bv.resize(m_numVisible);
|
||||
m_bh.resize(m_numHidden);
|
||||
}
|
||||
|
||||
void shuffle(double stdDev)
|
||||
@@ -70,6 +92,16 @@ public:
|
||||
uint32_t i, j;
|
||||
double kdev = stdDev*sqrt(12.0);
|
||||
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
{
|
||||
m_sigma(i) = 1; //kdev*Noise_Uniform(&m_noise);
|
||||
}
|
||||
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
{
|
||||
m_mean(i) = 0; //kdev*Noise_Uniform(&m_noise);
|
||||
}
|
||||
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
{
|
||||
m_bv(i) = 0; //kdev*Noise_Uniform(&m_noise);
|
||||
@@ -94,6 +126,8 @@ public:
|
||||
m_bv = rhs.m_bv;
|
||||
m_bh = rhs.m_bh;
|
||||
m_w = rhs.m_w;
|
||||
m_sigma = rhs.m_sigma;
|
||||
m_mean = rhs.m_mean;
|
||||
|
||||
return *this;
|
||||
}
|
||||
@@ -103,12 +137,22 @@ public:
|
||||
return m_w;
|
||||
}
|
||||
|
||||
VectorXd& visibleBias()
|
||||
RowVectorXd& visibleBias()
|
||||
{
|
||||
return m_bv;
|
||||
}
|
||||
|
||||
VectorXd& hiddenBias()
|
||||
RowVectorXd& sigma()
|
||||
{
|
||||
return m_sigma;
|
||||
}
|
||||
|
||||
RowVectorXd& mean()
|
||||
{
|
||||
return m_mean;
|
||||
}
|
||||
|
||||
RowVectorXd& hiddenBias()
|
||||
{
|
||||
return m_bh;
|
||||
}
|
||||
@@ -151,6 +195,16 @@ public:
|
||||
return m_numVisible;
|
||||
}
|
||||
|
||||
uint32_t getNumVisibleX()
|
||||
{
|
||||
return m_numVisibleX;
|
||||
}
|
||||
|
||||
uint32_t getNumVisibleY()
|
||||
{
|
||||
return m_numVisibleY;
|
||||
}
|
||||
|
||||
uint32_t getNumHidden()
|
||||
{
|
||||
return m_numHidden;
|
||||
@@ -166,7 +220,7 @@ public:
|
||||
return;
|
||||
|
||||
|
||||
fprintf(pFile, "%d %d\n", m_numVisible, m_numHidden);
|
||||
fprintf(pFile, "%d %d %d\n", m_numVisibleX, m_numVisibleY, m_numHidden);
|
||||
|
||||
uint32_t i, j;
|
||||
|
||||
@@ -192,7 +246,8 @@ public:
|
||||
|
||||
void load(const char *pFilename)
|
||||
{
|
||||
uint32_t numVisible;
|
||||
uint32_t numVisibleX;
|
||||
uint32_t numVisibleY;
|
||||
uint32_t numHidden;
|
||||
FILE *pFile;
|
||||
|
||||
@@ -201,25 +256,25 @@ public:
|
||||
if (!pFile)
|
||||
return;
|
||||
|
||||
fscanf(pFile, "%d %d\n", &numVisible, &numHidden);
|
||||
fscanf(pFile, "%d %d %d\n", &numVisibleX, &numVisibleY, &numHidden);
|
||||
|
||||
alloc(numVisible, numHidden);
|
||||
setUnits(numVisibleX, numVisibleY, numHidden);
|
||||
|
||||
uint32_t i, j;
|
||||
float v;
|
||||
for (i=0; i < numVisible; i++)
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
{
|
||||
fscanf(pFile, "%f", &v);
|
||||
m_bv(i) = v;
|
||||
}
|
||||
for (i=0; i < numHidden; i++)
|
||||
for (i=0; i < m_numHidden; i++)
|
||||
{
|
||||
fscanf(pFile, "%f", &v);
|
||||
m_bh(i) = v;
|
||||
}
|
||||
for (i=0; i < numVisible; i++)
|
||||
for (i=0; i < m_numVisible; i++)
|
||||
{
|
||||
for (j=0; j < numHidden; j++)
|
||||
for (j=0; j < m_numHidden; j++)
|
||||
{
|
||||
|
||||
fscanf(pFile, "%f", &v);
|
||||
@@ -232,27 +287,15 @@ public:
|
||||
|
||||
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;
|
||||
VectorXd m_bv;
|
||||
VectorXd m_bh;
|
||||
|
||||
void alloc(uint32_t numVisible, uint32_t numHidden)
|
||||
{
|
||||
if ((m_numVisible == numVisible) && (m_numHidden == numHidden))
|
||||
{
|
||||
return;
|
||||
}
|
||||
m_numVisible = numVisible;
|
||||
m_numHidden = numHidden;
|
||||
|
||||
m_w.resize(numVisible, numHidden);
|
||||
m_bv.resize(numVisible);
|
||||
m_bh.resize(numHidden);
|
||||
|
||||
shuffle(0);
|
||||
}
|
||||
RowVectorXd m_bv;
|
||||
RowVectorXd m_bh;
|
||||
RowVectorXd m_sigma;
|
||||
RowVectorXd m_mean;
|
||||
|
||||
void free()
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user