- added data normalization

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@45 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2014-11-12 13:08:20 +00:00
parent 6e7b8bc151
commit 549089a440
3 changed files with 62 additions and 0 deletions
+37
View File
@@ -52,6 +52,7 @@ public:
, m_doRaoBlackwell(false)
, m_useProbsForHiddenReconstruction(false)
, m_doSparse(false)
, m_doNormalizeData(false)
, m_numGibbs(1)
{
Noise_Init(&m_noise, 0x32727155);
@@ -260,6 +261,31 @@ public:
}
}
void normalizeData(MatrixXd &src, double mu, double sigma)
{
uint32_t i;
uint32_t size = src.rows();
double mean;
double stdDev;
for (i=0; i < size; i++)
{
mean = src.row(i).array().mean();
src.row(i).array() -= mean;
src.row(i).array() += mu;
}
for (i=0; i < size; i++)
{
src.row(i).array() *= src.row(i).array();
}
for (i=0; i < size; i++)
{
stdDev = sqrt(src.row(i).array().mean());
src.row(i).array() /= stdDev;
src.row(i).array() *= sigma;
}
}
void train2(const LayerArray<VisibleLayer> &vt, uint32_t numEpochs, uint32_t batchSize, double sigmaMin = 0.05)
{
uint32_t t, i;
@@ -293,6 +319,11 @@ public:
{
// t = (uint32_t)(0.5 + (vt.getSize()-1)*Noise_Uniform(&m_noise));
batch.row(i) = vt[i].states();
}
if (m_doNormalizeData)
{
normalizeData(batch, 0.0, m_sigma);
}
for (epoch=0; epoch < numEpochs; epoch++)
@@ -603,6 +634,11 @@ public:
m_doSparse = flag;
}
void setNormalizeData(bool flag)
{
m_doNormalizeData = flag;
}
void setNumGibbs(uint32_t value)
{
m_numGibbs = value;
@@ -646,6 +682,7 @@ private:
bool m_doRaoBlackwell;
bool m_useProbsForHiddenReconstruction;
bool m_doSparse;
bool m_doNormalizeData;
volatile bool m_doCancel;
uint32_t m_numGibbs;