Jingqiao-ucsc's picture
Upload WiSER private archive chunk: code_snapshots
13a1c10 verified
Raw
History Blame Contribute Delete
5.64 kB
"""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",
]