| """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) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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", |
| ] |
|
|