From c4aca832ba73a909c4a385e8fb50081418793cd4 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 3 Jul 2019 15:09:49 -0400 Subject: [PATCH] added single node distdataparallel --- pytorch_lightning/models/trainer.py | 45 +++++++++++++++++-- .../pt_overrides/override_data_parallel.py | 10 +++++ 2 files changed, 51 insertions(+), 4 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index c44afcec..b506a19f 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -5,8 +5,10 @@ from pytorch_lightning.root_module.memory import get_gpu_memory_map import traceback from pytorch_lightning.root_module.model_saving import TrainerIO from torch.optim.lr_scheduler import MultiStepLR -from pytorch_lightning.pt_overrides.override_data_parallel import LightningDataParallel +from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel import pdb +import torch.multiprocessing as mp +import torch.distributed as dist try: from apex import amp @@ -277,10 +279,45 @@ class Trainer(TrainerIO): # print model summary model.summarize() - # put on gpu if needed + # when GPU is called, spawn off a single worker for each gpu if self.on_gpu: - model = LightningDataParallel(model, device_ids=self.data_parallel_device_ids) - model.cuda(self.data_parallel_device_ids[0]) + rank = 0 + mp.spawn(self.__dp_train, nprocs=len(self.data_parallel_device_ids), args=(rank, model, self, )) + else: + self.__run_pretrain_routine(model) + + def __dp_train(self, gpu_nb, proc_rank, model, cluster_obj): + """ + Entry point into a DP thread + :param gpu_nb: + :param model: + :param cluster_obj: + :return: + """ + # TODO: pass in ip + ip = "127.0.0.1" + + # configure server + rank = proc_rank * len(self.data_parallel_device_ids) + gpu_nb + print(f"GPU: {gpu_nb} - Rank: {rank}") + world_size = self.cluster.per_experiment_nb_nodes * self.cluster.per_experiment_nb_gpus + dist.init_process_group("nccl", init_method=f'tcp://{ip}:12001', rank=rank, world_size=world_size) + + # copy model to each gpu + torch.cuda.set_device(gpu_nb) + model.cuda(gpu_nb) + model = LightningDistributedDataParallel(model, device_ids=[gpu_nb]) + + # continue training routine + self.__run_pretrain_routine(model) + + + def __run_pretrain_routine(self, model): + """ + Sanity check a few things before starting actual training + :param model: + :return: + """ # run tiny validation to make sure program won't crash during val _ = self.validate(model, self.val_dataloader, max_batches=self.nb_sanity_val_steps) diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index 4ca19a24..cf4f476e 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -1,4 +1,5 @@ from torch.nn import DataParallel +from torch.nn.parallel import DistributedDataParallel import threading import torch @@ -30,6 +31,15 @@ class LightningDataParallel(DataParallel): return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) +class LightningDistributedDataParallel(DistributedDataParallel): + """ + Override the forward call in lightning so it goes to training and validation step respectively + """ + + def parallel_apply(self, replicas, inputs, kwargs): + return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) + + def parallel_apply(modules, inputs, kwargs_tup=None, devices=None): r"""Applies each `module` in :attr:`modules` in parallel on arguments contained in :attr:`inputs` (positional) and :attr:`kwargs_tup` (keyword)