From cd0d29423691ac2096afc177a92a4af3a5c2e1b6 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 3 Jul 2019 16:44:18 -0400 Subject: [PATCH] clean up dead code --- .../pt_overrides/override_data_parallel.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index 78436a3d..1fe3d2da 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -45,6 +45,15 @@ class LightningDataParallel(DataParallel): def parallel_apply(self, replicas, inputs, kwargs): return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) + +class LightningDistributedDataParallel(DistributedDataParallel): + """ + Override the forward call in lightning so it goes to training and validation step respectively + """ + + def parallel_apply(self, replicas, inputs, kwargs): + return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) + def forward(self, *inputs, **kwargs): print('in forward') self._sync_params() @@ -74,15 +83,6 @@ class LightningDataParallel(DataParallel): return output -class LightningDistributedDataParallel(DistributedDataParallel): - """ - Override the forward call in lightning so it goes to training and validation step respectively - """ - - def parallel_apply(self, replicas, inputs, kwargs): - return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) - - def parallel_apply(modules, inputs, kwargs_tup=None, devices=None): r"""Applies each `module` in :attr:`modules` in parallel on arguments contained in :attr:`inputs` (positional) and :attr:`kwargs_tup` (keyword)