clean up dead code

This commit is contained in:
William Falcon
2019-07-03 16:43:05 -04:00
parent c10121c6ff
commit e8abbb1e75
@@ -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
return outputs