From 5bdad8a7b889f1a174ea22bbb988909227c31bb1 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 3 Jul 2019 16:46:14 -0400 Subject: [PATCH] clean up dead code --- .../pt_overrides/override_data_parallel.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index 1fe3d2da..ff4fce8c 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -55,19 +55,25 @@ class LightningDistributedDataParallel(DistributedDataParallel): return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) def forward(self, *inputs, **kwargs): - print('in forward') self._sync_params() if self.device_ids: inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids) if len(self.device_ids) == 1: - print('a') - output = self.module(*inputs[0], **kwargs[0]) + # -------------- + # LIGHTNING MOD + # -------------- + # normal + # output = self.module(*inputs[0], **kwargs[0]) + + # lightning + if self.module.training: + output = self.module.training_step(*inputs[0], **kwargs[0]) + else: + output = self.module.validation_step(*inputs[0], **kwargs[0]) else: - print('b') outputs = self.parallel_apply(self._module_copies[:len(inputs)], inputs, kwargs) output = self.gather(outputs, self.output_device) else: - print('c') output = self.module(*inputs, **kwargs) if torch.is_grad_enabled():