MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / to_ddp

Method to_ddp

detrsmpl/core/distributed_wrapper.py:67–86  ·  view source on GitHub ↗

Wrap models with separate MMDistributedDataParallel. It only wraps the modules with parameters.

(self, device_ids, dim, broadcast_buffers,
               find_unused_parameters, **kwargs)

Source from the content-addressed store, hash-verified

65 self.output_device = _get_device_index(device_ids[0], True)
66
67 def to_ddp(self, device_ids, dim, broadcast_buffers,
68 find_unused_parameters, **kwargs):
69 """Wrap models with separate MMDistributedDataParallel.
70
71 It only wraps the modules with parameters.
72 """
73 for name, module in self.module._modules.items():
74 if next(module.parameters(), None) is None:
75 module = module.cuda()
76 elif all(not p.requires_grad for p in module.parameters()):
77 module = module.cuda()
78 else:
79 module = MMDistributedDataParallel(
80 module.cuda(),
81 device_ids=device_ids,
82 dim=dim,
83 broadcast_buffers=broadcast_buffers,
84 find_unused_parameters=find_unused_parameters,
85 **kwargs)
86 self.module._modules[name] = module
87
88 def scatter(self, inputs, kwargs, device_ids):
89 """Scatter function.

Callers 1

__init__Method · 0.95

Calls 2

parametersMethod · 0.80
itemsMethod · 0.45

Tested by

no test coverage detected