diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index cf4f476e..78436a3d 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -1,5 +1,6 @@ from torch.nn import DataParallel from torch.nn.parallel import DistributedDataParallel +import itertools import threading import torch @@ -7,6 +8,20 @@ from torch.cuda._utils import _get_device_index import pdb +def _find_tensors(obj): + r""" + Recursively find all tensors contained in the specified object. + """ + if isinstance(obj, torch.Tensor): + return [obj] + if isinstance(obj, (list, tuple)): + return itertools.chain(*map(_find_tensors, obj)) + if isinstance(obj, dict): + return itertools.chain(*map(_find_tensors, obj.values())) + return [] + + + def get_a_var(obj): if isinstance(obj, torch.Tensor): return obj @@ -30,6 +45,34 @@ class LightningDataParallel(DataParallel): 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() + 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]) + 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(): + # We'll return the output object verbatim since it is a freeform + # object. We need to find any tensors in this object, though, + # because we need to figure out which parameters were used during + # this forward pass, to ensure we short circuit reduction for any + # unused parameters. Only if `find_unused_parameters` is set. + if self.find_unused_parameters: + self.reducer.prepare_for_backward(list(_find_tensors(output))) + else: + self.reducer.prepare_for_backward([]) + return output + class LightningDistributedDataParallel(DistributedDataParallel): """ @@ -112,4 +155,4 @@ def parallel_apply(modules, inputs, kwargs_tup=None, devices=None): if isinstance(output, Exception): raise output outputs.append(output) - return outputs \ No newline at end of file + return outputs