From 3ef7b07f5ea2c058bde0dd7b4994fc3b47d9a100 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Wed, 6 Nov 2019 06:12:51 +0000 Subject: [PATCH] - moved numEpochs and miniBatchSize to RBM::Params git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@606 b431acfa-c32f-4a4a-93f1-934dc6c82436 --- source/Rbm.cpp | 6 +++--- source/Rbm.hpp | 16 ++++++++++++++-- source/Stack.cpp | 8 ++++---- source/Stack.hpp | 4 ++-- source/main.cpp | 8 +++++--- test.prj | 4 +++- 6 files changed, 31 insertions(+), 15 deletions(-) diff --git a/source/Rbm.cpp b/source/Rbm.cpp index faaa6f4..608388a 100644 --- a/source/Rbm.cpp +++ b/source/Rbm.cpp @@ -57,7 +57,7 @@ Json::Value Rbm::toJson() const return rbm; } -void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener) +void Rbm::train(const arma::mat& batch, IListener* pListener) { Status status; size_t epoch; @@ -83,7 +83,7 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, while (status.trainingSizeRemain) { - size_t miniBatchSizeActual = std::min(miniBatchSize, status.trainingSizeRemain); + size_t miniBatchSizeActual = std::min(m_params.miniBatchSize, status.trainingSizeRemain); arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); status.trainingSizeRemain -= miniBatchSizeActual; batchRowIndex += miniBatchSizeActual; @@ -95,7 +95,7 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, arma::mat hid_state(miniBatchSizeActual, m_w.n_cols); arma::mat hid_probs(miniBatchSizeActual, m_w.n_cols); - for (epoch=0; epoch < numEpochs; epoch++) + for (epoch=0; epoch < m_params.numEpochs; epoch++) { // Create hidden layer base on training data diff --git a/source/Rbm.hpp b/source/Rbm.hpp index e7e0251..95ba41c 100644 --- a/source/Rbm.hpp +++ b/source/Rbm.hpp @@ -35,6 +35,8 @@ public: , gibbsDoSampleHidden(true) , doSampleBatch(false) , numGibbs(1) + , miniBatchSize(100) + , numEpochs(1000) { } @@ -50,6 +52,8 @@ public: params["gibbsDoSampleHidden"] = (int)gibbsDoSampleHidden; params["doSampleBatch"] = (int)doSampleBatch; params["numGibbs"] = (int)numGibbs; + params["miniBatchSize"] = (int)miniBatchSize; + params["numEpochs"] = (int)numEpochs; return params; } @@ -64,6 +68,8 @@ public: gibbsDoSampleHidden = params["gibbsDoSampleHidden"] == 1; doSampleBatch = params["doSampleBatch"] == 1; numGibbs = params["numGibbs"].asUInt(); + miniBatchSize = params["miniBatchSize"].asUInt(); + numEpochs = params["numEpochs"].asUInt(); } double weightDecay; @@ -74,6 +80,8 @@ public: bool gibbsDoSampleHidden; bool doSampleBatch; size_t numGibbs; + size_t miniBatchSize; + size_t numEpochs; }; struct Status @@ -114,7 +122,7 @@ public: virtual ~Rbm(); void weightsInit(double stddev, double mu=0.0); - void train(arma::mat const &batch, size_t miniBatchSize, size_t numEpochs, IListener *pListener); + void train(arma::mat const &batch, IListener *pListener); arma::mat toHiddenState(const arma::mat &visible) const; arma::mat toVisibleState(const arma::mat &hidden) const; @@ -127,8 +135,12 @@ public: Json::Value toJson() const; void fromJson(Json::Value params); + Params& params() + { + return m_params; + } private: - const Params m_params; + Params m_params; arma::mat sample(arma::mat const &src); static arma::mat probsLogistic(arma::mat const &src); void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5); diff --git a/source/Stack.cpp b/source/Stack.cpp index dec3890..17c4f46 100644 --- a/source/Stack.cpp +++ b/source/Stack.cpp @@ -153,17 +153,17 @@ bool Stack::saveWeights() return true; } -void Stack::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener) +void Stack::train(const arma::mat& batch, Rbm::IListener* pListener) { Layer *pLayer = m_pLayers; while(pLayer) { - train(pLayer->id(), batch, miniBatchSize, numEpochs, pListener); + train(pLayer->id(), batch, pListener); pLayer = pLayer->upper; } } -void Stack::train(size_t layerId, const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener) +void Stack::train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) { arma::mat thisBatch = batch; Layer *pLayer = m_pLayers; @@ -180,7 +180,7 @@ void Stack::train(size_t layerId, const arma::mat& batch, size_t miniBatchSize, if (pLayer) { std::cout << m_prjname << ": " << " Training of layer " << std::to_string(layerId) << std::endl; - pLayer->train(thisBatch, miniBatchSize, numEpochs, pListener); + pLayer->train(thisBatch, pListener); } } diff --git a/source/Stack.hpp b/source/Stack.hpp index 711e275..d85cf62 100644 --- a/source/Stack.hpp +++ b/source/Stack.hpp @@ -30,8 +30,8 @@ public: void addLayer(Layer *pLayer); Layer* getLayer(size_t layerId) const; - void train(size_t layerId, const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener); - void train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener); + void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener); + void train(const arma::mat& batch, Rbm::IListener* pListener); bool load(); bool save(); void weightsInit(double stddev); diff --git a/source/main.cpp b/source/main.cpp index 987c65c..962ac65 100644 --- a/source/main.cpp +++ b/source/main.cpp @@ -101,6 +101,9 @@ int main() stack.addLayer(layer); } + Layer *layer = stack.getLayer(0); + layer->params().learningRate = 0.02; + // Save project stack.save(); @@ -119,12 +122,11 @@ int main() #endif // Train stack - stack.train(batch, 100, 1000, &statusDisplay); - + stack.train(batch, &statusDisplay); + // Save weights stack.saveWeights(); - Layer *layer = stack.getLayer(0); arma::mat v = arma::randu(numTraining, layer->bv().n_elem); arma::mat h = layer->toHiddenProbs(v); arma::mat r = layer->toVisibleProbs(h); diff --git a/test.prj b/test.prj index 4a0f2f2..ab1b459 100644 --- a/test.prj +++ b/test.prj @@ -15,8 +15,10 @@ "doSampleBatch" : 0, "gibbsDoSampleHidden" : 1, "gibbsDoSampleVisible" : 0, - "learningRate" : 0.10000000000000001, + "learningRate" : 0.02, + "miniBatchSize" : 100, "momentum" : 0.5, + "numEpochs" : 1000, "numGibbs" : 1, "weightDecay" : 0 }