From 470f3e6d292ab0aedb3f2175a62ed8d86e950f2e Mon Sep 17 00:00:00 2001 From: William Falcon Date: Thu, 18 Jul 2019 11:08:48 -0400 Subject: [PATCH] added training router --- pytorch_lightning/models/trainer.py | 29 +++++++++++++++++------------ 1 file changed, 17 insertions(+), 12 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 224bb751..b9000ff4 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -17,7 +17,7 @@ import tqdm from pytorch_lightning.root_module.memory import get_gpu_memory_map from pytorch_lightning.root_module.model_saving import TrainerIO -from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel +from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel, LightningDataParallel try: @@ -152,6 +152,10 @@ class Trainer(TrainerIO): ''' warnings.warn(msg) + @property + def data_parallel(self): + return self.use_dp or self.use_ddp + def __determine_data_use_amount(self, train_percent_check, val_percent_check, test_percent_check, overfit_pct): """ Use less data for debugging purposes @@ -305,15 +309,15 @@ class Trainer(TrainerIO): # ----------------------------- def fit(self, model): - # when using gpus, first thing we do is spawn a new process between each worker - # multi-gpu and multi-nodes - if self.data_parallel: + # when using multi-node or DDP within a node start each module in a separate process + if self.use_ddp: self.experiment = self.experiment.get_meta_copy() - mp.spawn(self.dp_train, nprocs=len(self.data_parallel_device_ids), args=(model, )) + mp.spawn(self.ddp_train, nprocs=len(self.data_parallel_device_ids), args=(model, )) - # treat 1 gpu as a different case to avoid nccl bugs - elif self.data_parallel_device_ids is not None and len(self.data_parallel_device_ids) == 1: - self.single_gpu_train(model) + # 1 gpu or dp option triggers training using DP module + # easier to avoid NCCL issues + elif self.use_dp: + self.dp_train(model) else: # CHOOSE OPTIMIZER @@ -330,14 +334,15 @@ class Trainer(TrainerIO): self.__run_pretrain_routine(model) - def single_gpu_train(self, model): - # torch.cuda.set_device(0) - model.cuda(0) + def dp_train(self, model): # CHOOSE OPTIMIZER # filter out the weights that were done on gpu so we can load on good old cpus self.optimizers = model.configure_optimizers() + # attach model to DP + model = LightningDataParallel(model, device_ids=self.data_parallel_device_ids) + # run through amp wrapper if self.use_amp: # An example @@ -348,7 +353,7 @@ class Trainer(TrainerIO): self.__run_pretrain_routine(model) - def dp_train(self, gpu_nb, model): + def ddp_train(self, gpu_nb, model): """ Entry point into a DP thread :param gpu_nb: