mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
clean up dead code
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user