nur-dev's picture
Add files using upload-large-folder tool
e69b72a verified
Raw
History Blame Contribute Delete
3.08 kB
"""Pure DistributedDataParallel helpers driven by torchrun environment variables.
No DeepSpeed, no ZeRO, no CPU/parameter offload: this module only wraps the
standard ``torch.distributed`` NCCL process group and the ``torchrun`` launch
contract (``RANK``/``LOCAL_RANK``/``WORLD_SIZE``). When launched as a single
process (no ``WORLD_SIZE`` or ``WORLD_SIZE==1``) it degrades to plain
single-device execution, which is what the CPU unit tests exercise.
"""
from __future__ import annotations
import os
import random
from dataclasses import dataclass
import numpy as np
import torch
@dataclass(frozen=True, slots=True)
class DistInfo:
"""Resolved process-group topology for the current process."""
rank: int
local_rank: int
world_size: int
is_distributed: bool
@property
def is_main(self) -> bool:
return self.rank == 0
def setup_distributed(*, backend: str = "nccl") -> DistInfo:
"""Initialise the process group from torchrun env vars, if any.
Returns a :class:`DistInfo`. Safe to call when not launched by torchrun: it
reports a single-process topology and initialises nothing.
"""
world_size = int(os.environ.get("WORLD_SIZE", "1"))
if world_size <= 1:
return DistInfo(rank=0, local_rank=0, world_size=1, is_distributed=False)
import torch.distributed as dist
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
if torch.cuda.is_available():
torch.cuda.set_device(local_rank)
chosen_backend = backend
init_kwargs = {"device_id": torch.device("cuda", local_rank)}
else: # pragma: no cover - CPU multi-process is not used in practice.
chosen_backend = "gloo"
init_kwargs = {}
if not dist.is_initialized():
dist.init_process_group(backend=chosen_backend, **init_kwargs)
return DistInfo(
rank=dist.get_rank(),
local_rank=local_rank,
world_size=dist.get_world_size(),
is_distributed=True,
)
def cleanup_distributed(info: DistInfo) -> None:
"""Barrier + tear down the process group if one was created."""
if not info.is_distributed:
return
import torch.distributed as dist
if dist.is_initialized():
if torch.cuda.is_available():
dist.barrier(device_ids=[info.local_rank])
else:
dist.barrier()
dist.destroy_process_group()
def is_main_process(info: DistInfo) -> bool:
return info.rank == 0
def seed_everything(seed: int, *, rank: int = 0) -> None:
"""Seed Python, NumPy, and torch RNGs.
Model construction uses the base ``seed`` on every rank so initial weights
match before DDP broadcasts them; data-order randomness is decorrelated via
the rank offset. Full cuDNN determinism is intentionally not forced (it would
slow training); the seeds are recorded for reproducibility instead.
"""
effective = seed + rank
random.seed(effective)
np.random.seed(effective % (2**32))
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)