mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
fixes tpu data loader bug (#957)
* fixes tpu data loader bug * fixes tpu data loader bug
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user