git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@14 b431acfa-c32f-4a4a-93f1-934dc6c82436
416 lines
7.5 KiB
C++
416 lines
7.5 KiB
C++
/*
|
|
* Rbm.hpp
|
|
*
|
|
* Created on: 21.09.2014
|
|
* Author: jens
|
|
*/
|
|
|
|
#ifndef RBM_HPP_
|
|
#define RBM_HPP_
|
|
|
|
#include "Layer.hpp"
|
|
#include <cmath>
|
|
|
|
class VisibleLayer : public Layer
|
|
{
|
|
public:
|
|
VisibleLayer(uint32_t numUnits = 0)
|
|
: Layer(numUnits)
|
|
{
|
|
}
|
|
|
|
virtual ~VisibleLayer()
|
|
{
|
|
}
|
|
|
|
double getEnergy(const Weights &weights)
|
|
{
|
|
uint32_t i;
|
|
double energy = 0;
|
|
|
|
for (i=0; i < getNumUnits(); i++)
|
|
{
|
|
energy -= weights.getBiasVisible()[i] * getStates()[i];
|
|
}
|
|
return energy;
|
|
}
|
|
|
|
private:
|
|
double accum(const Layer &layer, const Weights &weights, uint32_t index) const
|
|
{
|
|
uint32_t i;
|
|
double sum = weights.getBiasVisible()[index];
|
|
const double *pStates = layer.getStates();
|
|
|
|
for (i=0; i < layer.getNumUnits(); i++)
|
|
{
|
|
sum += pStates[i] * weights.getWeights()[i][index];
|
|
}
|
|
|
|
return sum;
|
|
}
|
|
};
|
|
|
|
class HiddenLayer : public Layer
|
|
{
|
|
public:
|
|
HiddenLayer(uint32_t numUnits = 0)
|
|
: Layer(numUnits)
|
|
{
|
|
}
|
|
|
|
virtual ~HiddenLayer()
|
|
{
|
|
}
|
|
|
|
double getEnergy(const Weights &weights)
|
|
{
|
|
uint32_t i;
|
|
double energy = 0;
|
|
|
|
for (i=0; i < getNumUnits(); i++)
|
|
{
|
|
energy -= weights.getBiasHidden()[i] * getStates()[i];
|
|
}
|
|
return energy;
|
|
}
|
|
|
|
private:
|
|
double accum(const Layer &layer, const Weights &weights, uint32_t index) const
|
|
{
|
|
uint32_t i;
|
|
double sum = weights.getBiasVisible()[index];
|
|
const double *pStates = layer.getStates();
|
|
|
|
for (i=0; i < layer.getNumUnits(); i++)
|
|
{
|
|
sum += pStates[i] * weights.getWeights()[index][i];
|
|
}
|
|
|
|
return sum;
|
|
}
|
|
};
|
|
|
|
class Rbm
|
|
{
|
|
public:
|
|
Rbm(uint32_t numVisible, uint32_t numHidden)
|
|
: m_w(numVisible, numHidden)
|
|
, m_numVisible(numVisible)
|
|
, m_numHidden(numHidden)
|
|
, m_numTrainingPatterns(0)
|
|
, m_pVisibleTraining(nullptr)
|
|
{
|
|
Noise_Init(&m_noise, 0x32727155);
|
|
}
|
|
|
|
~Rbm()
|
|
{
|
|
freeTrainingPatterns();
|
|
Noise_Free(&m_noise);
|
|
}
|
|
|
|
void setTrainingInput(uint32_t index, const double *pValues)
|
|
{
|
|
if (!m_pVisibleTraining)
|
|
return;
|
|
|
|
if (index >= m_numTrainingPatterns)
|
|
return;
|
|
|
|
m_pVisibleTraining[index].setInput(pValues);
|
|
m_pVisibleTraining[index].statesAssignfromInput();
|
|
}
|
|
|
|
void setNumTrainingPatterns(uint32_t numTrainingPatterns)
|
|
{
|
|
m_numTrainingPatterns = numTrainingPatterns;
|
|
allocTrainingPatterns();
|
|
}
|
|
|
|
void weightsUpdate(VisibleLayer &v, VisibleLayer &vr, HiddenLayer &h, HiddenLayer &hr, double mu)
|
|
{
|
|
uint32_t i, j;
|
|
double dw;
|
|
|
|
// Update weights
|
|
for (i=0; i < m_numHidden; i++)
|
|
{
|
|
dw = 0;
|
|
for (j=0; j < m_numVisible; j++)
|
|
{
|
|
dw = v.getStates()[j] * h.getStates()[i];
|
|
m_w.getWeights()[i][j] += mu*dw;
|
|
}
|
|
}
|
|
|
|
for (i=0; i < m_numHidden; i++)
|
|
{
|
|
dw = 0;
|
|
for (j=0; j < m_numVisible; j++)
|
|
{
|
|
dw = vr.getStates()[j] * hr.getStates()[i];
|
|
m_w.getWeights()[i][j] -= mu*dw;
|
|
}
|
|
}
|
|
#if 1
|
|
for (i=0; i < m_numVisible; i++)
|
|
{
|
|
dw = v.getStates()[i] - vr.getStates()[i];
|
|
m_w.getBiasVisible()[i] += mu*dw;
|
|
}
|
|
|
|
for (i=0; i < m_numHidden; i++)
|
|
{
|
|
dw = h.getStates()[i] - hr.getStates()[i];
|
|
m_w.getBiasHidden()[i] += mu*dw;
|
|
}
|
|
#endif
|
|
}
|
|
|
|
void train(uint32_t numEpochs, double mu)
|
|
{
|
|
uint32_t epoch;
|
|
uint32_t trainingPatternIndex;
|
|
VisibleLayer *pV;
|
|
VisibleLayer vr(m_numVisible);
|
|
HiddenLayer h(m_numHidden);
|
|
HiddenLayer hr(m_numHidden);
|
|
const uint32_t monitorInterval = 100; // epochs
|
|
uint32_t monitorCount = monitorInterval; // epochs
|
|
|
|
for (epoch=0; epoch < numEpochs; epoch++)
|
|
{
|
|
trainingPatternIndex = (uint32_t)((m_numTrainingPatterns)*Noise_Uniform(&m_noise, 0.5));
|
|
|
|
if (trainingPatternIndex == m_numTrainingPatterns)
|
|
continue;
|
|
|
|
// Assign training data
|
|
pV = &m_pVisibleTraining[trainingPatternIndex];
|
|
|
|
// Create hidden layer base on training data
|
|
h.probsUpdate(*pV, m_w);
|
|
// h.statesAssignfromProbs();
|
|
h.statesUpdateStochastic();
|
|
|
|
// Create visible reconstruction (a fantasy...)
|
|
vr = *pV;
|
|
vr.probsUpdate(h, m_w);
|
|
vr.statesAssignfromProbs();
|
|
// vr.statesUpdateStochastic();
|
|
|
|
// Create hidden reconstruction
|
|
hr.probsUpdate(vr, m_w);
|
|
// hr.statesAssignfromProbs();
|
|
hr.statesUpdateStochastic();
|
|
|
|
// Update weights
|
|
weightsUpdate(*pV, vr, h, hr, mu);
|
|
|
|
if (!monitorCount)
|
|
{
|
|
monitorCount = monitorInterval;
|
|
// printf("Epoch #%d\n", epoch);
|
|
// prob();
|
|
}
|
|
monitorCount--;
|
|
}
|
|
}
|
|
|
|
double getEnergy(VisibleLayer &v, HiddenLayer &h)
|
|
{
|
|
uint32_t i, j;
|
|
double energy;
|
|
|
|
energy = -v.getEnergy(m_w) - h.getEnergy(m_w);
|
|
|
|
for (i=0; i < h.getNumUnits(); i++)
|
|
{
|
|
for (j=0; j < v.getNumUnits(); j++)
|
|
{
|
|
energy -= v.getStates()[j] * h.getStates()[i] * m_w.getWeights()[i][j];
|
|
}
|
|
}
|
|
return energy;
|
|
}
|
|
|
|
void prob()
|
|
{
|
|
uint32_t i, j;
|
|
double z;
|
|
double p;
|
|
|
|
HiddenLayer *h = new HiddenLayer[m_numTrainingPatterns];
|
|
|
|
// Create hidden layer activations based on training data
|
|
for (j=0; j < m_numTrainingPatterns; j++)
|
|
{
|
|
h[j].setNumUnits(m_numHidden);
|
|
h[j].probsUpdate(m_pVisibleTraining[j], m_w);
|
|
// h[j].statesAssignfromProbs();
|
|
h[j].statesUpdateStochastic();
|
|
}
|
|
|
|
printf("pi(t) = (pi^, v>)\n");
|
|
for (i=0; i < m_numHidden; i++)
|
|
{
|
|
for (j=0; j < m_numTrainingPatterns; j++)
|
|
{
|
|
p = h[j].getProbs()[i];
|
|
printf("%3.6f ", p);
|
|
}
|
|
printf("\n");
|
|
}
|
|
printf("\n");
|
|
|
|
printf("si(t) = (si^, v>)\n");
|
|
for (i=0; i < m_numHidden; i++)
|
|
{
|
|
for (j=0; j < m_numTrainingPatterns; j++)
|
|
{
|
|
p = h[j].getStates()[i];
|
|
printf("%3.6f ", p);
|
|
}
|
|
printf("\n");
|
|
}
|
|
printf("\n");
|
|
|
|
printf("p(v) = (t^, v>)\n");
|
|
for (i=0; i < m_numTrainingPatterns; i++)
|
|
{
|
|
z = 0;
|
|
for (j=0; j < m_numTrainingPatterns; j++)
|
|
{
|
|
z += exp(-getEnergy(m_pVisibleTraining[j], h[i]));
|
|
}
|
|
for (j=0; j < m_numTrainingPatterns; j++)
|
|
{
|
|
p = exp(-getEnergy(m_pVisibleTraining[j], h[i]))/z;
|
|
printf("%3.6f ", p);
|
|
}
|
|
printf("\n");
|
|
}
|
|
printf("\n");
|
|
|
|
// Reconstruct
|
|
for (i=0; i < m_numTrainingPatterns; i++)
|
|
{
|
|
m_pVisibleTraining[i].probsUpdate(h[i], m_w);
|
|
}
|
|
|
|
printf("A fantasy... (v^, t>)\n");
|
|
for (i=0; i < m_numVisible; i++)
|
|
{
|
|
for (j=0; j < m_numTrainingPatterns; j++)
|
|
{
|
|
p = m_pVisibleTraining[j].getProbs()[i];
|
|
printf("%3.6f ", p);
|
|
}
|
|
printf("\n");
|
|
}
|
|
|
|
delete [] h;
|
|
}
|
|
|
|
void toHidden(const double *pVisible)
|
|
{
|
|
double p;
|
|
uint32_t i;
|
|
|
|
VisibleLayer v(m_numVisible);
|
|
HiddenLayer h(m_numHidden);
|
|
|
|
v.setInput(pVisible);
|
|
v.statesAssignfromInput();
|
|
|
|
h.probsUpdate(v, m_w);
|
|
|
|
printf("pi(t) = (pi^, v>)\n");
|
|
for (i=0; i < m_numHidden; i++)
|
|
{
|
|
p = h.getProbs()[i];
|
|
printf("%3.6f\n", p);
|
|
}
|
|
printf("\n");
|
|
}
|
|
|
|
void toVisible(const double *pHidden)
|
|
{
|
|
double p;
|
|
uint32_t i;
|
|
|
|
VisibleLayer v(m_numVisible);
|
|
HiddenLayer h(m_numHidden);
|
|
|
|
h.setInput(pHidden);
|
|
h.statesAssignfromInput();
|
|
|
|
v.probsUpdate(h, m_w);
|
|
|
|
printf("pi(t) = (pi^, v>)\n");
|
|
for (i=0; i < m_numVisible; i++)
|
|
{
|
|
p = v.getProbs()[i];
|
|
printf("%3.6f\n", p);
|
|
}
|
|
printf("\n");
|
|
}
|
|
|
|
void weightsPrint()
|
|
{
|
|
m_w.print();
|
|
}
|
|
|
|
void weightsShuffle(double stdDev)
|
|
{
|
|
m_w.shuffle(stdDev);
|
|
}
|
|
|
|
private:
|
|
|
|
Weights m_w;
|
|
uint32_t m_numVisible;
|
|
uint32_t m_numHidden;
|
|
uint32_t m_numTrainingPatterns;
|
|
VisibleLayer *m_pVisibleTraining;
|
|
noise_gen_t m_noise;
|
|
|
|
void allocTrainingPatterns()
|
|
{
|
|
uint32_t i;
|
|
|
|
if (m_pVisibleTraining)
|
|
{
|
|
freeTrainingPatterns();
|
|
allocTrainingPatterns();
|
|
}
|
|
else
|
|
{
|
|
m_pVisibleTraining = new VisibleLayer[m_numTrainingPatterns];
|
|
for (i=0; i < m_numTrainingPatterns; i++)
|
|
{
|
|
m_pVisibleTraining[i].setNumUnits(m_numVisible);
|
|
}
|
|
}
|
|
}
|
|
|
|
void freeTrainingPatterns()
|
|
{
|
|
uint32_t i;
|
|
|
|
if (m_pVisibleTraining)
|
|
{
|
|
delete [] m_pVisibleTraining;
|
|
m_pVisibleTraining = nullptr;
|
|
}
|
|
m_numTrainingPatterns = 0;
|
|
}
|
|
|
|
};
|
|
|
|
|
|
|
|
|
|
#endif /* RBM_HPP_ */
|