Wrap models with separate MMDistributedDataParallel. It only wraps the modules with parameters.
(self, device_ids, dim, broadcast_buffers,
find_unused_parameters, **kwargs)
| 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. |
no test coverage detected