- improved progressIndicator

- fixed Stack::trainingData()
- fixed test::main

git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@638 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2019-11-08 06:38:09 +00:00
parent 215cfd11fa
commit 9ea6e48fff
5 changed files with 37 additions and 44 deletions
+1 -1
View File
@@ -845,7 +845,7 @@ bool MainComponent::onProgress(Rbm *pRbm, const Rbm::Status &status)
pComp->redrawReconstruction(); pComp->redrawReconstruction();
pComp->redrawWeights(); pComp->redrawWeights();
m_progressBarSlider->setValue(status.progress + 0.5); m_progressBarSlider->setValue(status.progress);
return !m_doStop; return !m_doStop;
} }
+14 -16
View File
@@ -63,11 +63,10 @@ Json::Value Rbm::toJson() const
void Rbm::train(const arma::mat& batch, IListener* pListener) void Rbm::train(const arma::mat& batch, IListener* pListener)
{ {
Status status; Status status;
size_t epoch; double dProgress = 100.0/(batch.n_rows*m_params.numEpochs);
size_t gibbs; double progress = 0;
int lastProgress = -100;
status.trainingSizeRemain = batch.n_rows; int batchRowIndex = 0;
size_t batchRowIndex = 0;
arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows)); arma::mat grad_bias_v(arma::zeros(1, m_w.n_rows));
arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols)); arma::mat grad_bias_h(arma::zeros(1, m_w.n_cols));
@@ -77,12 +76,14 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
arma::mat momentum_bias_h(arma::zeros(1, m_w.n_cols)); arma::mat momentum_bias_h(arma::zeros(1, m_w.n_cols));
arma::mat penalty_weights = arma::zeros(m_w.n_rows, m_w.n_cols); arma::mat penalty_weights = arma::zeros(m_w.n_rows, m_w.n_cols);
int trainingSizeRemain = batch.n_rows;
bool shouldAbort = false; bool shouldAbort = false;
while (status.trainingSizeRemain && !shouldAbort) while (trainingSizeRemain && !shouldAbort)
{ {
size_t miniBatchSizeActual = std::min(m_params.miniBatchSize, status.trainingSizeRemain); int miniBatchSizeActual = std::min(m_params.miniBatchSize, trainingSizeRemain);
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1); arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
status.trainingSizeRemain -= miniBatchSizeActual; trainingSizeRemain -= miniBatchSizeActual;
batchRowIndex += miniBatchSizeActual; batchRowIndex += miniBatchSizeActual;
double learning_rate = m_params.learningRate/miniBatchSizeActual; double learning_rate = m_params.learningRate/miniBatchSizeActual;
double weight_decay = m_params.weightDecay/miniBatchSizeActual; double weight_decay = m_params.weightDecay/miniBatchSizeActual;
@@ -92,13 +93,11 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
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);
double dProgress = 100.0/(batch.n_rows*m_params.numEpochs); for (int epoch=0; epoch < m_params.numEpochs; epoch++)
double lastProgress = -100.0;
for (epoch=0; epoch < m_params.numEpochs; epoch++)
{ {
if ((status.progress - lastProgress) >= 1.00) status.progress = (int)(progress + 0.5);
if (status.progress != lastProgress)
{ {
lastProgress = status.progress; lastProgress = status.progress;
if (pListener) if (pListener)
@@ -139,7 +138,7 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
grad_bias_v = sum(vis_state, 0); grad_bias_v = sum(vis_state, 0);
grad_bias_h = sum(hid_state, 0); grad_bias_h = sum(hid_state, 0);
for (gibbs=0; gibbs < m_params.numGibbs; gibbs++) for (int gibbs=0; gibbs < m_params.numGibbs; gibbs++)
{ {
// Create visible reconstruction (a fantasy...) given hid // Create visible reconstruction (a fantasy...) given hid
if (m_params.gibbsDoSampleHidden) if (m_params.gibbsDoSampleHidden)
@@ -181,11 +180,10 @@ void Rbm::train(const arma::mat& batch, IListener* pListener)
m_bh += learning_rate*momentum_bias_h; m_bh += learning_rate*momentum_bias_h;
m_w += learning_rate*momentum_weights; m_w += learning_rate*momentum_weights;
status.progress += dProgress*miniBatchSizeActual; progress += dProgress*miniBatchSizeActual;
} // Number of epochs } // Number of epochs
status.epoch = epoch;
arma::mat diffErr = miniBatch - vis_probs; arma::mat diffErr = miniBatch - vis_probs;
arma::mat diffErr_squared = diffErr % diffErr; arma::mat diffErr_squared = diffErr % diffErr;
status.err = accu(diffErr_squared)/diffErr_squared.n_elem; status.err = accu(diffErr_squared)/diffErr_squared.n_elem;
+8 -12
View File
@@ -68,9 +68,9 @@ public:
gibbsDoSampleVisible = params.get("gibbsDoSampleVisible", gibbsDoSampleVisible) == 1; gibbsDoSampleVisible = params.get("gibbsDoSampleVisible", gibbsDoSampleVisible) == 1;
gibbsDoSampleHidden = params.get("gibbsDoSampleHidden", gibbsDoSampleHidden) == 1; gibbsDoSampleHidden = params.get("gibbsDoSampleHidden", gibbsDoSampleHidden) == 1;
doSampleBatch = params.get("doSampleBatch", doSampleBatch) == 1; doSampleBatch = params.get("doSampleBatch", doSampleBatch) == 1;
numGibbs = params.get("numGibbs", (int)numGibbs).asUInt(); numGibbs = params.get("numGibbs", numGibbs).asUInt();
miniBatchSize = params.get("miniBatchSize", (int)miniBatchSize).asUInt(); miniBatchSize = params.get("miniBatchSize", miniBatchSize).asUInt();
numEpochs = params.get("numEpochs", (int)numEpochs).asUInt(); numEpochs = params.get("numEpochs", numEpochs).asUInt();
} }
double weightDecay; double weightDecay;
@@ -80,26 +80,22 @@ public:
bool gibbsDoSampleVisible; bool gibbsDoSampleVisible;
bool gibbsDoSampleHidden; bool gibbsDoSampleHidden;
bool doSampleBatch; bool doSampleBatch;
size_t numGibbs; int numGibbs;
size_t miniBatchSize; int miniBatchSize;
size_t numEpochs; int numEpochs;
}; };
struct Status struct Status
{ {
Status() Status()
: epoch(0) : progress(0)
, trainingSizeRemain(0)
, progress(0)
, err(-1.0) , err(-1.0)
, err_total(-1.0) , err_total(-1.0)
, L1(-1.0) , L1(-1.0)
, L2(-1.0) , L2(-1.0)
{ {
} }
size_t epoch; int progress;
size_t trainingSizeRemain;
double progress;
double err; double err;
double err_total; double err_total;
double L1; double L1;
+5 -4
View File
@@ -198,12 +198,13 @@ arma::mat& Stack::trainingData()
return m_trainingData; return m_trainingData;
} }
arma::mat Stack::trainingData(Layer* pLayer) arma::mat Stack::trainingData(Layer* pThatLayer)
{ {
arma::mat thisBatch = m_trainingData; arma::mat thisBatch = m_trainingData;
Layer *pThisLayer = m_pLayers; Layer *pLayer = m_pLayers;
while (pLayer) { while (pLayer)
if (pThisLayer->id() == pLayer->id()) {
if (pLayer->id() == pThatLayer->id())
{ {
break; break;
} }
+9 -11
View File
@@ -19,9 +19,7 @@ class RbmListener : public Rbm::IListener
bool onProgress(const Rbm::Status &status) bool onProgress(const Rbm::Status &status)
{ {
std::cout << "Progress : " << 100*status.progress << " %" << std::endl; std::cout << "Progress : " << status.progress << " %" << std::endl;
std::cout << "epoch : " << status.epoch << std::endl;
std::cout << "trainingSizeRemain: " << status.trainingSizeRemain << std::endl;
std::cout << "error (per mini batch) = " << status.err << std::endl; std::cout << "error (per mini batch) = " << status.err << std::endl;
std::cout << "error (total) = " << status.err_total << std::endl; std::cout << "error (total) = " << status.err_total << std::endl;
std::cout << "L1 = " << status.L1 << std::endl; std::cout << "L1 = " << status.L1 << std::endl;
@@ -84,15 +82,15 @@ int main()
RbmListener statusDisplay; RbmListener statusDisplay;
Stack stack(project); Stack stack(project);
arma::mat batch = stack.loadTraining(); stack.loadTraining();
printf("Loaded %d training samples\n", (int)batch.n_rows); printf("Loaded %d training samples\n", (int)stack.trainingData().n_rows);
stack.addTraining(batch, batch.row(1)); stack.addTraining(stack.trainingData().row(1));
printf("Loaded %d training samples\n", (int)batch.n_rows); printf("Loaded %d training samples\n", (int)stack.trainingData().n_rows);
stack.delTraining(batch, 0); stack.delTraining(0);
printf("Loaded %d training samples\n", (int)batch.n_rows); printf("Loaded %d training samples\n", (int)stack.trainingData().n_rows);
#if 1 #if 1
const int numLayers = 4; const int numLayers = 4;
@@ -127,13 +125,13 @@ int main()
#endif #endif
// Train stack // Train stack
stack.train(batch, &statusDisplay); stack.train(&statusDisplay);
// Save weights // Save weights
stack.saveWeights(); stack.saveWeights();
Layer *layer = stack.getLayer(0); Layer *layer = stack.getLayer(0);
arma::mat v = arma::randu(batch.n_rows, layer->bv().n_elem); arma::mat v = arma::randu(stack.trainingData().n_rows, 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);
return 0; return 0;