Download sacflow/utils/distributed.py from sathiiii/SACFlow: direct link, hf CLI and curl.
- Browser
- Download file 1.99 kB
-
https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/utils/distributed.py
- Command line
-
hf download hf://sathiiii/SACFlow/sacflow/utils/distributed.py
-
curl -L -o distributed.py https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/utils/distributed.py
1.99 kB
| 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 | |