MCPcopy Create free account
hub / github.com/OpenGVLab/HumanBench / __init__

Method __init__

PATH/core/distributed_utils.py:27–51  ·  view source on GitHub ↗
(self, module, sync=False, task_grp=None, share_backbone_group=None, \
            share_neck_group=None, share_decoder_group=None, ignore_bcast=None, \
            task_weight=None, task_size=None)

Source from the content-addressed store, hash-verified

25
26class DistModule(torch.nn.Module):
27 def __init__(self, module, sync=False, task_grp=None, share_backbone_group=None, \
28 share_neck_group=None, share_decoder_group=None, ignore_bcast=None, \
29 task_weight=None, task_size=None):
30 super(DistModule, self).__init__()
31 self.module = module
32 self.sync = sync
33 self.task_grp = task_grp
34 self.share_backbone_group = share_backbone_group
35 self.share_neck_group = share_neck_group
36 self.share_decoder_group = share_decoder_group
37 self.task_weight = task_weight
38 self.task_size = task_size
39
40 if not hasattr(torch.nn.Module, 'named_buffers'):
41 printlog('registering named_buffers for nn.Module at DistModule')
42 torch.nn.Module.named_buffers = named_buffers
43
44 broadcast_params_multitask(self, self.task_grp, self.share_backbone_group, \
45 self.share_neck_group, self.share_decoder_group, ignore_bcast)
46
47 assert sync, "Currently, only sync model is supported!"
48 if not sync:
49 self._grad_accs = {}
50 self._reduce_hooks = {}
51 self._register_hooks()
52
53 def forward(self, *inputs, **kwargs):
54 return self.module(*inputs, **kwargs)

Callers

nothing calls this directly

Calls 2

printlogFunction · 0.90

Tested by

no test coverage detected