mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-08-30 11:29:56 +08:00
finished dist (#911)
This commit is contained in:
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user