Rbm:🚋 better status indication
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@584 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+10
-8
@@ -59,11 +59,12 @@ Json::Value Rbm::toJson() const
|
|||||||
|
|
||||||
void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener)
|
void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs, IListener* pListener)
|
||||||
{
|
{
|
||||||
|
Status status;
|
||||||
size_t epoch;
|
size_t epoch;
|
||||||
size_t gibbs;
|
size_t gibbs;
|
||||||
|
|
||||||
size_t trainingSize = batch.n_rows;
|
size_t trainingSize = batch.n_rows;
|
||||||
size_t trainingSizeRemain = trainingSize;
|
status.trainingSizeRemain = trainingSize;
|
||||||
size_t batchRowIndex = 0;
|
size_t batchRowIndex = 0;
|
||||||
|
|
||||||
double dProgress = 1.0/trainingSize;
|
double dProgress = 1.0/trainingSize;
|
||||||
@@ -76,15 +77,16 @@ void Rbm::train(const arma::mat& batch, size_t miniBatchSize, size_t numEpochs,
|
|||||||
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);
|
||||||
|
|
||||||
Status status;
|
if (pListener)
|
||||||
|
|
||||||
status.progress = 0;
|
|
||||||
while (trainingSizeRemain)
|
|
||||||
{
|
{
|
||||||
status.trainingSizeRemain = trainingSizeRemain;
|
pListener->onProgress(status);
|
||||||
size_t miniBatchSizeActual = std::min(miniBatchSize, trainingSizeRemain);
|
}
|
||||||
|
|
||||||
|
while (status.trainingSizeRemain)
|
||||||
|
{
|
||||||
|
size_t miniBatchSizeActual = std::min(miniBatchSize, status.trainingSizeRemain);
|
||||||
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
arma::mat miniBatch = batch.rows(batchRowIndex, batchRowIndex+miniBatchSizeActual-1);
|
||||||
trainingSizeRemain -= miniBatchSizeActual;
|
status.trainingSizeRemain -= miniBatchSizeActual;
|
||||||
batchRowIndex += miniBatchSizeActual;
|
batchRowIndex += miniBatchSizeActual;
|
||||||
double learning_rate = m_params.learningRate/std::min(miniBatchSizeActual, trainingSize);
|
double learning_rate = m_params.learningRate/std::min(miniBatchSizeActual, trainingSize);
|
||||||
double weight_decay = m_params.weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
double weight_decay = m_params.weightDecay/std::min(miniBatchSizeActual, trainingSize);
|
||||||
|
|||||||
@@ -81,6 +81,16 @@ public:
|
|||||||
|
|
||||||
struct Status
|
struct Status
|
||||||
{
|
{
|
||||||
|
Status()
|
||||||
|
: epoch(0)
|
||||||
|
, trainingSizeRemain(0)
|
||||||
|
, progress(0)
|
||||||
|
, err(-1.0)
|
||||||
|
, err_total(1-0)
|
||||||
|
, L1(-1.0)
|
||||||
|
, L2(-1.0)
|
||||||
|
{
|
||||||
|
}
|
||||||
size_t epoch;
|
size_t epoch;
|
||||||
size_t trainingSizeRemain;
|
size_t trainingSizeRemain;
|
||||||
double progress;
|
double progress;
|
||||||
|
|||||||
Reference in New Issue
Block a user