From 4dccddd50f31c7af2e0f065f86a3114c8e25b16f Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Wed, 15 Jun 2016 21:30:28 +0000 Subject: [PATCH] [RBM] - added getHiddenBatch() and getVisibleBatch() - cleaned up git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@292 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- Source/Rbm.hpp | 78 ++++++++++++++++++++++++----------------- Source/RbmComponent.cpp | 18 +++++----- Source/RbmComponent.h | 5 +-- 3 files changed, 56 insertions(+), 45 deletions(-) diff --git a/Source/Rbm.hpp b/Source/Rbm.hpp index a0aa2a4..1dea022 100644 --- a/Source/Rbm.hpp +++ b/Source/Rbm.hpp @@ -38,7 +38,7 @@ public: class Rbm { public: - Rbm(Weights &weights, const LayerArray &batch, RbmListener *pListener = nullptr) + Rbm(Weights &weights, const MatrixXd &batch, RbmListener *pListener = nullptr) : m_w(weights) , m_batch(batch) , m_variableSigma(weights.getNumVisible()) @@ -224,13 +224,15 @@ public: uint32_t t, i; uint32_t epoch; uint32_t gibbs; + + size_t batchSize = m_batch.rows(); + double dProgress = 1.0/numEpochs; - double kTrain = 1.0/m_batch.getSize(); + double kTrain = 1.0/batchSize; - size_t batchSize = m_batch.getSize(); - MatrixXd v(batchSize, m_w.getNumVisible()); - MatrixXd h(batchSize, m_w.getNumHidden()); + m_v.resize(batchSize, m_w.getNumVisible()); + m_h.resize(batchSize, m_w.getNumHidden()); MatrixXd sumBiasV(1, m_w.getNumVisible()); MatrixXd sumBiasH(1, m_w.getNumHidden()); @@ -245,7 +247,7 @@ public: m_progress = 0; - MatrixXd batch = m_batch.data(); + MatrixXd batch = m_batch; if (m_doNormalizeData) { @@ -262,62 +264,62 @@ public: double err; // Create hidden layer base on training data - toHiddenBatch(h, batch); - probsLogistic(h); + toHiddenBatch(m_h, batch); + probsLogistic(m_h); if (!m_doRaoBlackwell) { - sample(h); + sample(m_h); } // Update weights (positive phase) sumBiasV = batch.colwise().sum(); if (!m_doSparse) { - sumBiasH = h.colwise().sum(); + sumBiasH = m_h.colwise().sum(); } - sumWeights = batch.transpose() * h; + sumWeights = batch.transpose() * m_h; diffErr = batch; for (gibbs=0; gibbs < m_numGibbs; gibbs++) { - sample(h); + sample(m_h); // Create visible reconstruction (a fantasy...) given h - toVisibleBatch(v, h); + toVisibleBatch(m_v, m_h); if (m_useVisibleGaussian) { if (!m_useProbsForHiddenReconstruction) { - sampleGaussian(v, m_variableSigma.replicate(batchSize, 1)); + sampleGaussian(m_v, m_variableSigma.replicate(batchSize, 1)); } } else { - probsLogistic(v, m_variableSigma.replicate(batchSize, 1)); + probsLogistic(m_v, m_variableSigma.replicate(batchSize, 1)); if (!m_useProbsForHiddenReconstruction) { - sample(v); + sample(m_v); } } // Create hidden reconstruction given v - toHiddenBatch(h, v); - probsLogistic(h); + toHiddenBatch(m_h, m_v); + probsLogistic(m_h); } if (!m_doRaoBlackwell) { - sample(h); + sample(m_h); } // Update weights (negative phase) - sumBiasV -= v.colwise().sum(); + sumBiasV -= m_v.colwise().sum(); if (!m_doSparse) { - sumBiasH -= h.colwise().sum(); + sumBiasH -= m_h.colwise().sum(); } - sumWeights -= v.transpose() * h; - diffErr -= v; + sumWeights -= m_v.transpose() * m_h; + diffErr -= m_v; deltaWeights = m_momentum*deltaWeights + m_muWeights*(kTrain*sumWeights - m_weightDecay*m_w.weights()); m_w.weights() += deltaWeights; @@ -327,12 +329,12 @@ public: if (m_doSparse) { - h = v * m_w.weights(); - h += m_w.hiddenBias().replicate(batchSize, 1); - probsLogistic(h); + m_h = m_v * m_w.weights(); + m_h += m_w.hiddenBias().replicate(batchSize, 1); + probsLogistic(m_h); sumBiasH.fill(m_sparsity); - sumBiasH -= h.colwise().mean(); + sumBiasH -= m_h.colwise().mean(); deltaBiasH = m_momentum*deltaBiasH + m_muSparsity*sumBiasH; @@ -468,9 +470,9 @@ public: void setDoLearnVariance(bool flag) { m_doLearnVariance = flag; - if (m_doLearnVariance && m_batch.getSize()) + if (m_doLearnVariance && m_batch.rows()) { - m_variableSigma = calcSigma(m_batch.data()); + m_variableSigma = calcSigma(m_batch); } else { @@ -503,9 +505,21 @@ public: m_momentum = value; } + MatrixXd const& getHiddenBatch() + { + return m_h; + } + + MatrixXd const& getVisibleBatch() + { + return m_v; + } + private: Weights &m_w; - LayerArray const &m_batch; + MatrixXd const &m_batch; + MatrixXd m_v; + MatrixXd m_h; RowVectorXd m_variableSigma; double m_constantSigma; RbmListener *m_pListener; @@ -529,13 +543,13 @@ private: void toHiddenBatch(MatrixXd &h, MatrixXd const &v) { h = v * m_w.weights(); - h += m_w.hiddenBias().replicate(m_batch.getSize(), 1); + h += m_w.hiddenBias().replicate(m_batch.rows(), 1); } void toVisibleBatch(MatrixXd &v, MatrixXd const &h) { v = h * m_w.weights().transpose(); - v += m_w.visibleBias().replicate(m_batch.getSize(), 1); + v += m_w.visibleBias().replicate(m_batch.rows(), 1); } }; diff --git a/Source/RbmComponent.cpp b/Source/RbmComponent.cpp index 2ef85a6..77659e7 100644 --- a/Source/RbmComponent.cpp +++ b/Source/RbmComponent.cpp @@ -36,18 +36,18 @@ RbmComponent::RbmComponent (Weights &weights, RbmComponentListener &listener) { //[UserPreSize] - m_pRbm = new Rbm(m_weights, m_layers, this); - m_vNumX = m_weights.getNumVisibleX(); - m_vNumY = m_weights.getNumVisibleY(); - m_hNum = m_weights.getNumHidden(); + m_pRbm = new Rbm(m_weights, m_layers.data(), this); + size_t vNumX = m_weights.getNumVisibleX(); + size_t vNumY = m_weights.getNumVisibleY(); + size_t hNum = m_weights.getNumHidden(); - addAndMakeVisible (DrawTraining = new DrawComponent (m_vNumX, m_vNumY)); + addAndMakeVisible (DrawTraining = new DrawComponent (vNumX, vNumY)); DrawTraining->setListener(this); - addAndMakeVisible (DrawReconstruction = new DrawComponent (m_vNumX, m_vNumY)); - addAndMakeVisible (DrawWeights = new DrawComponent (m_vNumX, m_vNumY)); - addAndMakeVisible (DrawVars = new DrawComponent (m_vNumX, m_vNumY)); - addAndMakeVisible (DrawHidden = new DrawComponent (m_hNum, 1)); + addAndMakeVisible (DrawReconstruction = new DrawComponent (vNumX, vNumY)); + addAndMakeVisible (DrawWeights = new DrawComponent (vNumX, vNumY)); + addAndMakeVisible (DrawVars = new DrawComponent (vNumX, vNumY)); + addAndMakeVisible (DrawHidden = new DrawComponent (hNum, 1)); DrawHidden->setListener(this); //[/UserPreSize] diff --git a/Source/RbmComponent.h b/Source/RbmComponent.h index 3d846fb..1d03223 100644 --- a/Source/RbmComponent.h +++ b/Source/RbmComponent.h @@ -109,7 +109,7 @@ public: DrawReconstruction->getData() = DrawTraining->getData(); DrawReconstruction->DrawData(); - m_layers.add(DrawTraining->getData(), m_vNumX*m_vNumY); + m_layers.add(DrawTraining->getData(), m_weights.getNumVisible()); m_listener.onLayerSizeChanged(m_layers.getSize()); } @@ -177,9 +177,6 @@ private: ScopedPointer DrawWeights; ScopedPointer DrawVars; ScopedPointer DrawHidden; - uint32_t m_vNumX; - uint32_t m_vNumY; - uint32_t m_hNum; LayerArray m_layers; size_t m_currWeightIndexToDraw; size_t m_currTrainingIndexToDraw;