Sync ``module``'s parameters and buffers state. Syncs ``module``'s parameters and buffers state so that all ranks contain the same module state across all ranks. Note that this API assumes that all parameter shapes are consistent before running the synchronization. This can be
(
module: nn.Module,
process_group: dist.ProcessGroup,
broadcast_bucket_size: int,
src: int,
params_and_buffers_to_ignore: Container[str],
broadcast_buffers: bool = True,
)
| 264 | |
| 265 | |
| 266 | def _sync_module_states( |
| 267 | module: nn.Module, |
| 268 | process_group: dist.ProcessGroup, |
| 269 | broadcast_bucket_size: int, |
| 270 | src: int, |
| 271 | params_and_buffers_to_ignore: Container[str], |
| 272 | broadcast_buffers: bool = True, |
| 273 | ) -> None: |
| 274 | """ |
| 275 | Sync ``module``'s parameters and buffers state. |
| 276 | |
| 277 | Syncs ``module``'s parameters and buffers state so that all ranks contain |
| 278 | the same module state across all ranks. Note that this API assumes that all |
| 279 | parameter shapes are consistent before running the synchronization. This can |
| 280 | be checked with ``_verify_param_shape_across_processes``. |
| 281 | """ |
| 282 | module_states: List[torch.Tensor] = [] |
| 283 | for name, param in module.named_parameters(): |
| 284 | if name not in params_and_buffers_to_ignore: |
| 285 | module_states.append(param.detach()) |
| 286 | |
| 287 | if broadcast_buffers: |
| 288 | for name, buffer in module.named_buffers(): |
| 289 | if name not in params_and_buffers_to_ignore: |
| 290 | module_states.append(buffer.detach()) |
| 291 | |
| 292 | _sync_params_and_buffers(process_group, module_states, broadcast_bucket_size, src) |
| 293 | |
| 294 | |
| 295 | def _sync_params_and_buffers( |
searching dependent graphs…