scaled batch size

This commit is contained in:
William Falcon
2019-07-08 19:48:22 -04:00
parent 25dbd7a936
commit f2c1f0221e
@@ -162,15 +162,13 @@ class LightningTemplateModel(LightningModule):
train_sampler = None
batch_size = self.hparams.batch_size
try:
if self.on_gpu:
train_sampler = DistributedSampler(dataset, num_replicas=self.trainer.world_size, rank=self.trainer.proc_rank)
# try:
# if self.on_gpu:
train_sampler = DistributedSampler(dataset, num_replicas=self.trainer.world_size, rank=self.trainer.proc_rank)
batch_size = batch_size // self.trainer.world_size # scale batch size
# scale batch size
batch_size = batch_size // self.trainer.world_size
except Exception as e:
pass
# except Exception as e:
# pass
should_shuffle = train_sampler is None
loader = DataLoader(