"""DDP-aware logging + env helpers. Usage: import os, torch.distributed as dist from wiser.utils.ddp_logger import ( setup_ddp, is_rank0, get_rank, get_world_size, get_local_rank, setup_logger, log, barrier, ) setup_ddp() logger = setup_logger(out_dir) log("starting training") # goes to both rank-specific and combined log Every rank writes to `{out_dir}/artifacts/rank_{R}.log`; rank 0 additionally mirrors to `{out_dir}/artifacts/train.log` (combined feed so you can watch progress without opening 8 files). """ from __future__ import annotations import logging import os import pathlib import sys import traceback from datetime import datetime _LOGGER: logging.Logger | None = None _LOG_FMT = "%(asctime)s [rank %(rank_s)s] [pid %(process)d] %(levelname)s %(message)s" _LOG_DATEFMT = "%H:%M:%S" def _env_int(key: str, default: int = 0) -> int: try: return int(os.environ.get(key, default)) except (TypeError, ValueError): return default def get_rank() -> int: """Global rank. 0 if not running under torchrun.""" return _env_int("RANK", 0) def get_local_rank() -> int: """Local (per-node) rank. Used to select cuda device.""" return _env_int("LOCAL_RANK", 0) def get_world_size() -> int: """Total number of ranks. 1 if not running under torchrun.""" return max(1, _env_int("WORLD_SIZE", 1)) def is_rank0() -> bool: return get_rank() == 0 def is_distributed() -> bool: """True iff torchrun-style env vars indicate multi-rank run.""" return get_world_size() > 1 def setup_ddp() -> tuple[int, int, int]: """Initialize torch.distributed process group if running under torchrun. Returns (rank, local_rank, world_size). Safe to call when world_size=1 (becomes no-op). """ rank = get_rank() local_rank = get_local_rank() world_size = get_world_size() if world_size > 1: import torch import torch.distributed as dist if not dist.is_initialized(): backend = "nccl" if torch.cuda.is_available() else "gloo" dist.init_process_group(backend=backend) torch.cuda.set_device(local_rank) return rank, local_rank, world_size def barrier() -> None: """All-rank synchronization point. No-op in single-rank mode.""" if is_distributed(): import torch.distributed as dist if dist.is_initialized(): dist.barrier() def cleanup_ddp() -> None: if is_distributed(): import torch.distributed as dist if dist.is_initialized(): dist.destroy_process_group() class _RankFilter(logging.Filter): def __init__(self, rank: int) -> None: super().__init__() self._rank = rank def filter(self, record: logging.LogRecord) -> bool: record.rank_s = str(self._rank) return True def setup_logger(out_dir: pathlib.Path | str) -> logging.Logger: """Set up a root logger with per-rank file + (rank 0) combined file + stderr. Idempotent — safe to call multiple times (returns the same logger). """ global _LOGGER if _LOGGER is not None: return _LOGGER rank = get_rank() out_path = pathlib.Path(out_dir) / "artifacts" out_path.mkdir(parents=True, exist_ok=True) logger = logging.getLogger("a06_9.train") logger.setLevel(logging.INFO) logger.handlers.clear() fmt = logging.Formatter(_LOG_FMT, datefmt=_LOG_DATEFMT) rank_filter = _RankFilter(rank) # Per-rank file (exclusive writer). rank_file = out_path / f"rank_{rank}.log" rank_handler = logging.FileHandler(rank_file, mode="a", encoding="utf-8") rank_handler.setFormatter(fmt) rank_handler.addFilter(rank_filter) logger.addHandler(rank_handler) # Rank 0: also write a combined 'train.log' so tail -f is easy. if rank == 0: combined_file = out_path / "train.log" combined_handler = logging.FileHandler(combined_file, mode="a", encoding="utf-8") combined_handler.setFormatter(fmt) combined_handler.addFilter(rank_filter) logger.addHandler(combined_handler) # Stderr mirror on rank 0 only (so other ranks don't spam console). if rank == 0: stream_handler = logging.StreamHandler(sys.stderr) stream_handler.setFormatter(fmt) stream_handler.addFilter(rank_filter) logger.addHandler(stream_handler) logger.propagate = False _LOGGER = logger log( f"logger initialised: rank={rank}/{get_world_size()} " f"local_rank={get_local_rank()} out_dir={out_path}" ) return logger def log(message: str, level: int = logging.INFO) -> None: """Log a message with the per-rank formatted prefix. Falls back to print() if logger hasn't been initialised yet (early boot). """ if _LOGGER is None: print(f"[{datetime.now().strftime(_LOG_DATEFMT)}] [rank {get_rank()}] {message}", file=sys.stderr, flush=True) return _LOGGER.log(level, message) def log_exception(message: str) -> None: """Log an exception with full traceback on the current rank.""" tb = traceback.format_exc() log(f"{message}\n{tb}", level=logging.ERROR) def log_rank0(message: str, level: int = logging.INFO) -> None: """Log only on rank 0. Useful for summary messages.""" if is_rank0(): log(message, level=level) __all__ = [ "setup_ddp", "setup_logger", "log", "log_rank0", "log_exception", "barrier", "cleanup_ddp", "get_rank", "get_local_rank", "get_world_size", "is_rank0", "is_distributed", ]