(launcher: str, backend: str = 'nccl', **kwargs)
| 50 | |
| 51 | |
| 52 | def init_dist(launcher: str, backend: str = 'nccl', **kwargs) -> None: |
| 53 | if mp.get_start_method(allow_none=True) is None: |
| 54 | mp.set_start_method('spawn') |
| 55 | if launcher == 'pytorch': |
| 56 | _init_dist_pytorch(backend, **kwargs) |
| 57 | elif launcher == 'mpi': |
| 58 | _init_dist_mpi(backend, **kwargs) |
| 59 | elif launcher == 'slurm': |
| 60 | _init_dist_slurm(backend, **kwargs) |
| 61 | else: |
| 62 | raise ValueError(f'Invalid launcher type: {launcher}') |
| 63 | |
| 64 | |
| 65 | def _init_dist_pytorch(backend: str, **kwargs) -> None: |
searching dependent graphs…