mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
added dist sampler exception
This commit is contained in:
@@ -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
|
||||
# -----------------------------
|
||||
|
||||
Reference in New Issue
Block a user