File size: 1,992 Bytes
d65ae7d | 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 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 | from __future__ import annotations
import os
from datetime import timedelta
import torch
import torch.distributed as dist
def is_dist_avail_and_initialized() -> bool:
return dist.is_available() and dist.is_initialized()
def get_rank() -> int:
if not is_dist_avail_and_initialized():
return 0
return dist.get_rank()
def get_world_size() -> int:
if not is_dist_avail_and_initialized():
return 1
return dist.get_world_size()
def is_main_process() -> bool:
return get_rank() == 0
def init_distributed(backend: str = "nccl") -> torch.device:
"""Initialize distributed training.
Validation on 3D medical volumes can take longer than PyTorch's default
10 minute NCCL/RCCL watchdog timeout if some ranks are waiting at a
collective. We therefore set a longer timeout by default. The value can be
overridden with DIST_TIMEOUT_MINUTES.
"""
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
local_rank = int(os.environ.get("LOCAL_RANK", 0))
torch.cuda.set_device(local_rank)
timeout_min = int(os.environ.get("DIST_TIMEOUT_MINUTES", "180"))
dist.init_process_group(
backend=backend,
init_method="env://",
timeout=timedelta(minutes=timeout_min),
)
device = torch.device("cuda", local_rank)
else:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
return device
def barrier():
if is_dist_avail_and_initialized():
if torch.cuda.is_available():
dist.barrier(device_ids=[torch.cuda.current_device()])
else:
dist.barrier()
def cleanup():
if is_dist_avail_and_initialized():
dist.destroy_process_group()
def reduce_mean(tensor: torch.Tensor) -> torch.Tensor:
if not is_dist_avail_and_initialized():
return tensor
rt = tensor.detach().clone()
dist.all_reduce(rt, op=dist.ReduceOp.SUM)
rt /= get_world_size()
return rt
|