- 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:
+3
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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);
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user