added dist sampler exception

This commit is contained in:
William Falcon
2019-07-08 19:39:59 -04:00
parent 96314cbf46
commit 3e2dde1680
+16
View File
@@ -12,6 +12,7 @@ import torch.distributed as dist
import os
import subprocess
from time import sleep
from torch.utils.data.distributed import DistributedSampler
try:
@@ -270,6 +271,21 @@ class Trainer(TrainerIO):
self.test_dataloader = model.test_dataloader
self.val_dataloader = model.val_dataloader
if self.on_gpu and type(self.tng_dataloader.sampler) is not DistributedSampler:
msg = '''
when using multiple gpus and multiple nodes you must pass a DistributedSampler to DataLoader(sampler).
ie: this:
dataset = myDataset()
dataloader = Dataloader(dataset)
becomes:
dataset = myDataset()
dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset)
dataloader = Dataloader(dataset, sampler=dist_sampler)
'''
raise Exception(msg)
# -----------------------------
# MODEL TRAINING
# -----------------------------