diff --git a/pytorch_lightning/trainer/evaluation_loop.py b/pytorch_lightning/trainer/evaluation_loop.py index bca62836..a62d318a 100644 --- a/pytorch_lightning/trainer/evaluation_loop.py +++ b/pytorch_lightning/trainer/evaluation_loop.py @@ -257,7 +257,6 @@ class TrainerEvaluationLoopMixin(ABC): dataloader = dataloader.per_device_loader(device) for batch_idx, batch in enumerate(dataloader): - if batch is None: # pragma: no cover continue diff --git a/pytorch_lightning/trainer/training_loop.py b/pytorch_lightning/trainer/training_loop.py index d2be9894..ef056f88 100644 --- a/pytorch_lightning/trainer/training_loop.py +++ b/pytorch_lightning/trainer/training_loop.py @@ -456,15 +456,18 @@ class TrainerTrainLoopMixin(ABC): if self.reload_dataloaders_every_epoch: self.reset_train_dataloader(self.get_model()) + # track local dataloader so TPU can wrap each epoch + train_dataloader = self.train_dataloader + # on TPU we have to wrap it under the ParallelLoader if self.use_tpu: device = xm.xla_device() - self.train_dataloader = xla_pl.ParallelLoader(self.train_dataloader, [device]) - self.train_dataloader = self.train_dataloader.per_device_loader(device) + train_dataloader = xla_pl.ParallelLoader(train_dataloader, [device]) + train_dataloader = train_dataloader.per_device_loader(device) # run epoch for batch_idx, batch in self.profiler.profile_iterable( - enumerate(self.train_dataloader), "get_train_batch" + enumerate(train_dataloader), "get_train_batch" ): # stop epoch if we limited the number of training batches if batch_idx >= self.num_training_batches: diff --git a/tests/models/utils.py b/tests/models/utils.py index 75be02ef..30a11a7a 100644 --- a/tests/models/utils.py +++ b/tests/models/utils.py @@ -188,7 +188,7 @@ def run_prediction(dataloader, trained_model, dp=False, min_acc=0.50): acc = torch.tensor(acc) acc = acc.item() - assert acc > min_acc, f'this model is expected to get > {min_acc} in test set (it got {acc})' + assert acc >= min_acc, f'this model is expected to get > {min_acc} in test set (it got {acc})' def assert_ok_model_acc(trainer, key='test_acc', thr=0.4):