| 54 | |
| 55 | |
| 56 | def initialize_global_process_group(timeout_second=36000): |
| 57 | torch.distributed.init_process_group( |
| 58 | get_nccl_backend(), |
| 59 | timeout=timedelta(seconds=timeout_second), |
| 60 | init_method=os.environ.get("DIST_INIT_METHOD", None), |
| 61 | ) |
| 62 | local_rank = int(os.environ["LOCAL_RANK"]) |
| 63 | rank = int(os.environ["RANK"]) |
| 64 | world_size = int(os.environ["WORLD_SIZE"]) |
| 65 | |
| 66 | if torch.distributed.is_initialized(): |
| 67 | get_torch_device().set_device(local_rank) |
| 68 | return local_rank, rank, world_size |
| 69 | |
| 70 | |
| 71 | def destroy_global_process_group(): |