added single node distdataparallel

This commit is contained in:
William Falcon
2019-07-03 15:09:49 -04:00
parent 30e2fc6c4b
commit c4aca832ba
2 changed files with 51 additions and 4 deletions
+41 -4
View File
@@ -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)
@@ -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)