diff --git a/pytorch_lightning/trainer/data_loading.py b/pytorch_lightning/trainer/data_loading.py index 496debbf..ee545128 100644 --- a/pytorch_lightning/trainer/data_loading.py +++ b/pytorch_lightning/trainer/data_loading.py @@ -25,6 +25,15 @@ try: except ImportError: APEX_AVAILABLE = False +try: + import torch_xla + import torch_xla.core.xla_model as xm + import torch_xla.distributed.xla_multiprocessing as xmp + + XLA_AVAILABLE = True +except ImportError: + XLA_AVAILABLE = False + class TrainerDataLoadingMixin(ABC): @@ -217,12 +226,15 @@ class TrainerDataLoadingMixin(ABC): # on TPUs load each dataloader only on process 0 # this will trigger the data downloads - if self.use_tpu: + if self.use_tpu and XLA_AVAILABLE: if self.tpu_local_core_rank == 0: self.get_train_dataloader() self.get_test_dataloaders() self.get_val_dataloaders() + # wait for all processes to catch up + torch_xla.core.xla_model.rendezvous() + # support IterableDataset for train data self.is_iterable_train_dataloader = ( EXIST_ITER_DATASET and isinstance(self.get_train_dataloader().dataset, IterableDataset)) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 6cd87b17..e2d76943 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -983,6 +983,11 @@ class Trainer(TrainerIOMixin, if self.use_ddp or self.use_ddp2: dist.barrier() + # wait for all models to restore weights + if self.on_tpu and XLA_AVAILABLE: + # wait for all processes to catch up + torch_xla.core.xla_model.rendezvous() + # set up checkpoint callback self.configure_checkpoint_callback() diff --git a/pytorch_lightning/trainer/training_io.py b/pytorch_lightning/trainer/training_io.py index 4d439a35..1bc14966 100644 --- a/pytorch_lightning/trainer/training_io.py +++ b/pytorch_lightning/trainer/training_io.py @@ -106,6 +106,15 @@ from pytorch_lightning.overrides.data_parallel import ( LightningDataParallel, ) +try: + import torch_xla + import torch_xla.core.xla_model as xm + import torch_xla.distributed.xla_multiprocessing as xmp + + XLA_AVAILABLE = True +except ImportError: + XLA_AVAILABLE = False + class TrainerIOMixin(ABC): @@ -125,6 +134,7 @@ class TrainerIOMixin(ABC): self.early_stop_callback = None self.lr_schedulers = None self.optimizers = None + self.on_tpu = None self.num_training_batches = None self.accumulate_grad_batches = None @@ -170,6 +180,11 @@ class TrainerIOMixin(ABC): # wait for all processes to catch up dist.barrier() + # wait for all models to restore weights + if self.on_tpu and XLA_AVAILABLE: + # wait for all processes to catch up + torch_xla.core.xla_model.rendezvous() + # clear cache after restore if self.on_gpu: torch.cuda.empty_cache()