File size: 1,471 Bytes
3ce19a2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 | from datetime import timedelta
import os
import torch
import torch.distributed as dist
def init_distributed_mode():
if "MASTER_ADDR" in os.environ and "MASTER_PORT" in os.environ:
dist_url = f"tcp://{os.environ['MASTER_ADDR']}:{os.environ['MASTER_PORT']}"
else:
dist_url = "env://"
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
local_rank = int(os.environ.get("LOCAL_RANK", 0))
elif "SLURM_NODEID" in os.environ:
gpus_per_node = torch.cuda.device_count()
node_id = int(os.environ["SLURM_NODEID"])
local_rank = int(os.environ["SLURM_LOCALID"])
rank = node_id * gpus_per_node + local_rank
world_size = int(os.environ["SLURM_NTASKS"])
else:
raise RuntimeError("Distributed environment not properly set.")
torch.cuda.set_device(local_rank)
dist.init_process_group(
backend="nccl",
init_method=dist_url,
world_size=world_size,
rank=rank,
timeout=timedelta(seconds=4800),
)
dist.barrier()
def is_dist_avail_and_initialized():
return dist.is_available() and dist.is_initialized()
def get_world_size():
return dist.get_world_size() if is_dist_avail_and_initialized() else 1
def get_rank():
return dist.get_rank() if is_dist_avail_and_initialized() else 0
def is_main_process():
return get_rank() == 0
|