From 4b30ef6480f984420cab56da6466a377f1ddef22 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Tue, 12 May 2020 00:09:48 -0400 Subject: [PATCH] Device (#1790) * added self.device * added docs --- docs/source/multi_gpu.rst | 4 +++- pytorch_lightning/core/lightning.py | 3 +++ pytorch_lightning/trainer/distrib_data_parallel.py | 1 + pytorch_lightning/trainer/distrib_parts.py | 5 +++++ pytorch_lightning/trainer/trainer.py | 1 + 5 files changed, 13 insertions(+), 1 deletion(-) diff --git a/docs/source/multi_gpu.rst b/docs/source/multi_gpu.rst index 8688cd33..9e32b2b0 100644 --- a/docs/source/multi_gpu.rst +++ b/docs/source/multi_gpu.rst @@ -46,7 +46,9 @@ This will make your code scale to any arbitrary number of GPUs or TPUs with Ligh # with lightning def forward(self, x): z = torch.Tensor(2, 3) - z = z.type_as(x) + z = z.type_as(x, device=self.device) + +Every LightningModule knows what device it is on. You can access that reference via `self.device`. Remove samplers ^^^^^^^^^^^^^^^ diff --git a/pytorch_lightning/core/lightning.py b/pytorch_lightning/core/lightning.py index 8638a512..32dd13a7 100644 --- a/pytorch_lightning/core/lightning.py +++ b/pytorch_lightning/core/lightning.py @@ -72,6 +72,9 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks): self.hparams = None + #: device reference + self.device = None + def print(self, *args, **kwargs) -> None: r""" Prints only from process 0. Use this in any distributed mode to log only once. diff --git a/pytorch_lightning/trainer/distrib_data_parallel.py b/pytorch_lightning/trainer/distrib_data_parallel.py index 8651dd5c..4bf0c7ff 100644 --- a/pytorch_lightning/trainer/distrib_data_parallel.py +++ b/pytorch_lightning/trainer/distrib_data_parallel.py @@ -344,6 +344,7 @@ class TrainerDDPMixin(ABC): # copy model to each gpu if self.on_gpu: self.root_gpu = process_idx + self.device = torch.device('cuda', self.root_gpu) torch.cuda.set_device(self.root_gpu) model.cuda(self.root_gpu) diff --git a/pytorch_lightning/trainer/distrib_parts.py b/pytorch_lightning/trainer/distrib_parts.py index bcd0c072..8c54faf5 100644 --- a/pytorch_lightning/trainer/distrib_parts.py +++ b/pytorch_lightning/trainer/distrib_parts.py @@ -432,6 +432,7 @@ class TrainerDPMixin(ABC): m.use_tpu = self.use_tpu m.tpu_local_core_rank = self.tpu_local_core_rank m.tpu_global_core_rank = self.tpu_global_core_rank + m.device = self.device def transfer_batch_to_tpu(self, batch): return self.__transfer_data_to_device(batch, device='tpu') @@ -483,6 +484,7 @@ class TrainerDPMixin(ABC): def single_gpu_train(self, model): model.cuda(self.root_gpu) + self.device = torch.device('cuda', self.root_gpu) # CHOOSE OPTIMIZER # allow for lr schedulers as well @@ -499,6 +501,7 @@ class TrainerDPMixin(ABC): def tpu_train(self, tpu_core_idx, model): # put model on tpu model.to(xm.xla_device()) + self.device = xm.xla_device() # get the appropriate tpu ranks self.tpu_local_core_rank = xm.get_local_ordinal() @@ -536,6 +539,7 @@ class TrainerDPMixin(ABC): self.optimizers, self.lr_schedulers, self.optimizer_frequencies = self.init_optimizers(model) model.cuda(self.root_gpu) + self.device = torch.device('cuda', self.root_gpu) # hack forward to do autocast for the user model_autocast_original_forward = model.forward @@ -575,6 +579,7 @@ class TrainerDPMixin(ABC): assert self.root_gpu == hvd.local_rank() torch.cuda.set_device(self.root_gpu) model.cuda(self.root_gpu) + self.device = torch.device('cuda', self.root_gpu) # avoid duplicating progress bar if hvd.rank() != 0 and self.progress_bar_callback is not None: diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index f34aa488..743d8618 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -462,6 +462,7 @@ class Trainer( # distributed backend choice self.distributed_backend = distributed_backend self.set_distributed_mode(distributed_backend) + self.device = torch.device('cpu') # override dist backend when using tpus if self.on_tpu: