- 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:
+49
-16
@@ -10,7 +10,7 @@
|
||||
|
||||
#ifndef LAYER_HPP
|
||||
#define LAYER_HPP
|
||||
#include <cstdint>
|
||||
#include <stdint.h>
|
||||
#include "noise.h"
|
||||
#include "Weights.hpp"
|
||||
|
||||
@@ -21,9 +21,8 @@ public:
|
||||
: m_numUnits(numUnits)
|
||||
, m_pProbs(nullptr)
|
||||
, m_pStates(nullptr)
|
||||
, m_pStatesInit(pStatesInit)
|
||||
{
|
||||
setNumUnits(numUnits);
|
||||
setNumUnits(numUnits, pStatesInit);
|
||||
Noise_Init(&m_noise, 0x12345677);
|
||||
}
|
||||
|
||||
@@ -33,7 +32,7 @@ public:
|
||||
Noise_Free(&m_noise);
|
||||
}
|
||||
|
||||
void setNumUnits(uint32_t numUnits)
|
||||
void setNumUnits(uint32_t numUnits, const double *pStatesInit = nullptr)
|
||||
{
|
||||
if (m_numUnits)
|
||||
{
|
||||
@@ -43,23 +42,17 @@ public:
|
||||
m_numUnits = numUnits;
|
||||
if (m_numUnits)
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
m_pProbs = new double[m_numUnits];
|
||||
m_pStates = new double[m_numUnits];
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
probsInit(0);
|
||||
if (pStatesInit)
|
||||
{
|
||||
m_pProbs[i] = 0.0;
|
||||
memcpy(m_pStates, pStatesInit, m_numUnits*sizeof(double));
|
||||
}
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
else
|
||||
{
|
||||
m_pStates[i] = 0.0;
|
||||
}
|
||||
if (m_pStatesInit)
|
||||
{
|
||||
memcpy(m_pStates, m_pStatesInit, m_numUnits*sizeof(double));
|
||||
statesInit(0);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -72,6 +65,37 @@ public:
|
||||
return *this;
|
||||
}
|
||||
|
||||
Layer& operator+= (const Layer &rhs)
|
||||
{
|
||||
uint32_t i;
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pStates[i] += rhs.m_pStates[i];
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
void probsInit(double value) const
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pProbs[i] = value;
|
||||
}
|
||||
}
|
||||
|
||||
void statesInit(double value) const
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pStates[i] = value;
|
||||
}
|
||||
}
|
||||
|
||||
void probsUpdate(const Layer &layer, const Weights &weights) const
|
||||
{
|
||||
uint32_t i;
|
||||
@@ -82,6 +106,16 @@ public:
|
||||
}
|
||||
}
|
||||
|
||||
void statesScale(double kscale) const
|
||||
{
|
||||
uint32_t i;
|
||||
|
||||
for (i=0; i < m_numUnits; i++)
|
||||
{
|
||||
m_pStates[i] *= kscale;
|
||||
}
|
||||
}
|
||||
|
||||
void statesAssignfromProbs()
|
||||
{
|
||||
memcpy(m_pStates, m_pProbs, m_numUnits*sizeof(double));
|
||||
@@ -128,7 +162,6 @@ private:
|
||||
protected:
|
||||
double *m_pProbs;
|
||||
double *m_pStates;
|
||||
const double *m_pStatesInit;
|
||||
|
||||
virtual double accum(const Layer &layer, const Weights &weights, uint32_t index) const = 0;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user