Get the data parallel group the caller rank belongs to.
()
| 317 | |
| 318 | |
| 319 | def _get_data_parallel_group(): |
| 320 | """Get the data parallel group the caller rank belongs to.""" |
| 321 | assert dist.is_initialized(), \ |
| 322 | 'dist is not initialized' |
| 323 | global mpu |
| 324 | if mpu is not None: |
| 325 | return mpu.get_data_parallel_group() |
| 326 | # Return the clone of dist world group |
| 327 | return _clone_world_group() |
| 328 | |
| 329 | |
| 330 | def _get_broadcast_src_rank(): |
no test coverage detected