scaled batch size

This commit is contained in:
William Falcon
2019-07-08 19:45:52 -04:00
parent 971a6c4184
commit 25dbd7a936
@@ -163,7 +163,7 @@ class LightningTemplateModel(LightningModule):
batch_size = self.hparams.batch_size
try:
if self.hparams.nb_gpu_nodes > 1:
if self.on_gpu:
train_sampler = DistributedSampler(dataset, num_replicas=self.trainer.world_size, rank=self.trainer.proc_rank)
# scale batch size