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