From 9f8ab7c29e1fc42d1ab472a311e54cd052aa6582 Mon Sep 17 00:00:00 2001 From: Vadim Bereznyuk Date: Wed, 30 Oct 2019 19:13:40 +0300 Subject: [PATCH] Fixed total number of batches (#439) * Fixed total number of batches * Fixed flake8 warning * Update train_loop_mixin.py * Update train_loop_mixin.py --- pytorch_lightning/trainer/train_loop_mixin.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/pytorch_lightning/trainer/train_loop_mixin.py b/pytorch_lightning/trainer/train_loop_mixin.py index 8cebe7a1..2683ae66 100644 --- a/pytorch_lightning/trainer/train_loop_mixin.py +++ b/pytorch_lightning/trainer/train_loop_mixin.py @@ -23,7 +23,15 @@ class TrainerTrainLoopMixin(object): # update training progress in trainer and model model.current_epoch = epoch_nb self.current_epoch = epoch_nb - self.total_batches = self.nb_training_batches + self.nb_val_batches + + # val can be checked multiple times in epoch + is_val_epoch = (self.current_epoch + 1) % self.check_val_every_n_epoch == 0 + val_checks_per_epoch = self.nb_training_batches // self.val_check_batch + val_checks_per_epoch = val_checks_per_epoch if is_val_epoch else 0 + + # total batches includes multiple val checks + self.total_batches = (self.nb_training_batches + + self.nb_val_batches * val_checks_per_epoch) self.batch_loss_value = 0 # accumulated grads # limit the number of batches to 1 in fast_dev_run