moved sampler

This commit is contained in:
William Falcon
2019-07-08 18:02:41 -04:00
parent 14d1329655
commit bd2d1ddc07
2 changed files with 15 additions and 16 deletions
@@ -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
+1 -13
View File
@@ -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
# -----------------------------