- moved numEpochs and miniBatchSize to RBM::Params

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@606 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-06 06:12:51 +00:00
parent c834a86f0f
commit 3ef7b07f5e
6 changed files with 31 additions and 15 deletions
+3 -3
View File
@@ -57,7 +57,7 @@ Json::Value Rbm::toJson() const
return rbm; 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; Status status;
size_t epoch; size_t epoch;
@@ -83,7 +83,7 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
while (status.trainingSizeRemain) 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); arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
status.trainingSizeRemain -= miniBatchSizeActual; status.trainingSizeRemain -= miniBatchSizeActual;
batchRowIndex += 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_state(miniBatchSizeActual, m_w.n_cols);
arma::mat hid_probs(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 // Create hidden layer base on training data
+14 -2
View File
@@ -35,6 +35,8 @@ public:
, gibbsDoSampleHidden(true) , gibbsDoSampleHidden(true)
, doSampleBatch(false) , doSampleBatch(false)
, numGibbs(1) , numGibbs(1)
, miniBatchSize(100)
, numEpochs(1000)
{ {
} }
@@ -50,6 +52,8 @@ public:
params["gibbsDoSampleHidden"] = (int)gibbsDoSampleHidden; params["gibbsDoSampleHidden"] = (int)gibbsDoSampleHidden;
params["doSampleBatch"] = (int)doSampleBatch; params["doSampleBatch"] = (int)doSampleBatch;
params["numGibbs"] = (int)numGibbs; params["numGibbs"] = (int)numGibbs;
params["miniBatchSize"] = (int)miniBatchSize;
params["numEpochs"] = (int)numEpochs;
return params; return params;
} }
@@ -64,6 +68,8 @@ public:
gibbsDoSampleHidden = params["gibbsDoSampleHidden"] == 1; gibbsDoSampleHidden = params["gibbsDoSampleHidden"] == 1;
doSampleBatch = params["doSampleBatch"] == 1; doSampleBatch = params["doSampleBatch"] == 1;
numGibbs = params["numGibbs"].asUInt(); numGibbs = params["numGibbs"].asUInt();
miniBatchSize = params["miniBatchSize"].asUInt();
numEpochs = params["numEpochs"].asUInt();
} }
double weightDecay; double weightDecay;
@@ -74,6 +80,8 @@ public:
bool gibbsDoSampleHidden; bool gibbsDoSampleHidden;
bool doSampleBatch; bool doSampleBatch;
size_t numGibbs; size_t numGibbs;
size_t miniBatchSize;
size_t numEpochs;
}; };
struct Status struct Status
@@ -114,7 +122,7 @@ public:
virtual ~Rbm(); virtual ~Rbm();
void weightsInit(double stddev, double mu=0.0); 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 toHiddenState(const arma::mat &visible) const;
arma::mat toVisibleState(const arma::mat &hidden) const; arma::mat toVisibleState(const arma::mat &hidden) const;
@@ -127,8 +135,12 @@ public:
Json::Value toJson() const; Json::Value toJson() const;
void fromJson(Json::Value params); void fromJson(Json::Value params);
Params& params()
{
return m_params;
}
private: private:
const Params m_params; Params m_params;
arma::mat sample(arma::mat const &src); arma::mat sample(arma::mat const &src);
static arma::mat probsLogistic(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); void uniform(arma::mat &srcDst, double stdDev=1.0, double mu=0.5);
+4 -4
View File
@@ -153,17 +153,17 @@ bool Stack::saveWeights()
return true; 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; Layer *pLayer = m_pLayers;
while(pLayer) while(pLayer)
{ {
train(pLayer->id(), batch, miniBatchSize, numEpochs, pListener); train(pLayer->id(), batch, pListener);
pLayer = pLayer->upper; 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; arma::mat thisBatch = batch;
Layer *pLayer = m_pLayers; Layer *pLayer = m_pLayers;
@@ -180,7 +180,7 @@ void Stack::train(size_t layerId, const arma::mat& batch, size_t miniBatchSize,
if (pLayer) if (pLayer)
{ {
std::cout << m_prjname << ": " << " Training of layer " << std::to_string(layerId) << std::endl; std::cout << m_prjname << ": " << " Training of layer " << std::to_string(layerId) << std::endl;
pLayer->train(thisBatch, miniBatchSize, numEpochs, pListener); pLayer->train(thisBatch, pListener);
} }
} }
+2 -2
View File
@@ -30,8 +30,8 @@ public:
void addLayer(Layer *pLayer); void addLayer(Layer *pLayer);
Layer* getLayer(size_t layerId) const; 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(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener);
void train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, Rbm::IListener* pListener); void train(const arma::mat& batch, Rbm::IListener* pListener);
bool load(); bool load();
bool save(); bool save();
void weightsInit(double stddev); void weightsInit(double stddev);
+5 -3
View File
@@ -101,6 +101,9 @@ int main()
stack.addLayer(layer); stack.addLayer(layer);
} }
Layer *layer = stack.getLayer(0);
layer->params().learningRate = 0.02;
// Save project // Save project
stack.save(); stack.save();
@@ -119,12 +122,11 @@ int main()
#endif #endif
// Train stack // Train stack
stack.train(batch, 100, 1000, &statusDisplay); stack.train(batch, &statusDisplay);
// Save weights // Save weights
stack.saveWeights(); stack.saveWeights();
Layer *layer = stack.getLayer(0);
arma::mat v = arma::randu(numTraining, layer->bv().n_elem); arma::mat v = arma::randu(numTraining, layer->bv().n_elem);
arma::mat h = layer->toHiddenProbs(v); arma::mat h = layer->toHiddenProbs(v);
arma::mat r = layer->toVisibleProbs(h); arma::mat r = layer->toVisibleProbs(h);
+3 -1
View File
@@ -15,8 +15,10 @@
"doSampleBatch" : 0, "doSampleBatch" : 0,
"gibbsDoSampleHidden" : 1, "gibbsDoSampleHidden" : 1,
"gibbsDoSampleVisible" : 0, "gibbsDoSampleVisible" : 0,
"learningRate" : 0.10000000000000001, "learningRate" : 0.02,
"miniBatchSize" : 100,
"momentum" : 0.5, "momentum" : 0.5,
"numEpochs" : 1000,
"numGibbs" : 1, "numGibbs" : 1,
"weightDecay" : 0 "weightDecay" : 0
} }