| 237 | |
| 238 | |
| 239 | def world_info_from_env(): |
| 240 | local_rank = 0 |
| 241 | for v in ('LOCAL_RANK', 'MPI_LOCALRANKID', 'SLURM_LOCALID', 'OMPI_COMM_WORLD_LOCAL_RANK'): |
| 242 | if v in os.environ: |
| 243 | local_rank = int(os.environ[v]) |
| 244 | break |
| 245 | global_rank = 0 |
| 246 | for v in ('RANK', 'PMI_RANK', 'SLURM_PROCID', 'OMPI_COMM_WORLD_RANK'): |
| 247 | if v in os.environ: |
| 248 | global_rank = int(os.environ[v]) |
| 249 | break |
| 250 | world_size = 1 |
| 251 | for v in ('WORLD_SIZE', 'PMI_SIZE', 'SLURM_NTASKS', 'OMPI_COMM_WORLD_SIZE'): |
| 252 | if v in os.environ: |
| 253 | world_size = int(os.environ[v]) |
| 254 | break |
| 255 | |
| 256 | return local_rank, global_rank, world_size |
| 257 | |
| 258 | |
| 259 | def setup_distributed(backend="nccl", port=None): |