diff --git a/pytorch_lightning/trainer/data_loading.py b/pytorch_lightning/trainer/data_loading.py index 92fb73e2..b3e15024 100644 --- a/pytorch_lightning/trainer/data_loading.py +++ b/pytorch_lightning/trainer/data_loading.py @@ -133,8 +133,8 @@ class TrainerDataLoadingMixin(ABC): world_size = { 'ddp': self.num_nodes * self.num_processes, 'ddp2': self.num_nodes, + 'ddp_cpu': self.num_processes * self.num_nodes } - import pdb; pdb.set_trace() sampler = DistributedSampler( dataloader.dataset, num_replicas=world_size.get(self.distributed_backend, 0),