diff --git a/examples/new_project_templates/lightning_module_template.py b/examples/new_project_templates/lightning_module_template.py index b1ea9d9b..4d6a402f 100644 --- a/examples/new_project_templates/lightning_module_template.py +++ b/examples/new_project_templates/lightning_module_template.py @@ -7,6 +7,8 @@ import torch import torch.nn.functional as F from test_tube import HyperOptArgumentParser from torch import optim +from torch.utils.data import DataLoader +from torch.utils.data.distributed import DistributedSampler from pytorch_lightning.root_module.root_module import LightningModule @@ -154,13 +156,22 @@ class LightningTemplateModel(LightningModule): def __dataloader(self, train): # init data generators transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) - dataset = MNIST(root=self.hparams.data_root, train=train, transform=transform, download=True) - loader = torch.utils.data.DataLoader( + # when using multi-node we need to add the datasampler + try: + if self.hparams.nb_gpu_nodes > 1: + train_sampler = DistributedSampler(dataset, num_replicas=self.trainer.world_size, rank=self.trainer.proc_rank) + print('using sampler') + except Exception as e: + print('no sampler') + train_sampler = None + + loader = DataLoader( dataset=dataset, batch_size=self.hparams.batch_size, - shuffle=True + shuffle=True, + sampler=train_sampler ) return loader diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index fdef6be9..bb63e35b 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -12,7 +12,7 @@ import torch.distributed as dist import os import subprocess from time import sleep -from torch.utils.data.distributed import DistributedSampler + try: from apex import amp @@ -268,18 +268,6 @@ class Trainer(TrainerIO): self.test_dataloader = model.test_dataloader self.val_dataloader = model.val_dataloader - # when distributed data parallel, we need to distribute the dataset to each node - if self.nb_gpu_nodes > 1: - self.tng_dataloader = self.__distribute_dataloader(self.tng_dataloader) - self.test_dataloader = self.__distribute_dataloader(self.test_dataloader) - self.val_dataloader = self.__distribute_dataloader(self.val_dataloader) - - def __distribute_dataloader(self, dataloader): - dataset = dataloader.dataset - dataset = DistributedSampler(dataset, num_replicas=self.world_size, rank=self.proc_rank) - dataloader.dataset = dataset - return dataloader - # ----------------------------- # MODEL TRAINING # -----------------------------