- 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;
}
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
+14 -2
View File
@@ -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);
+4 -4
View File
@@ -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);
}
}
+2 -2
View File
@@ -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);
+5 -3
View File
@@ -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);
+3 -1
View File
@@ -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
}