Spaces:
Sleeping
Sleeping
| import os, torch | |
| from torch.distributed import init_process_group, destroy_process_group | |
| def ddp_env(): | |
| """ | |
| Reads ranks from torchrun env vars. | |
| """ | |
| if "RANK" not in os.environ: | |
| os.environ["RANK"] = "0" | |
| os.environ["WORLD_SIZE"] = "1" | |
| os.environ["LOCAL_RANK"] = "0" | |
| rank = int(os.environ["RANK"]) | |
| world_size = int(os.environ["WORLD_SIZE"]) | |
| local_rank = int(os.environ["LOCAL_RANK"]) | |
| return rank, world_size, local_rank | |
| def ddp_setup_from_env(): | |
| rank, world_size, local_rank = ddp_env() | |
| use_ddp = world_size > 1 | |
| if torch.cuda.is_available(): | |
| torch.cuda.set_device(local_rank) | |
| if use_ddp: | |
| init_process_group(backend="nccl") # rank/world_size inferred from env | |
| return rank, world_size, local_rank, use_ddp |