fixes tpu data loader bug (#957)

* fixes tpu data loader bug

* fixes tpu data loader bug
This commit is contained in:
William Falcon
2020-02-26 19:29:03 -05:00
committed by GitHub
parent b2e9607362
commit f86dd55145
3 changed files with 7 additions and 5 deletions
@@ -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
+6 -3
View File
@@ -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:
+1 -1
View File
@@ -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):