""" Optional memory diagnostics for long-running training (CPU RSS, worker tree, CUDA, tracemalloc). YAML (all optional except interval to enable): memory_debug_interval: 2000 # 0 = disabled memory_debug_tracemalloc: false # true = on RSS spike, diff Python allocs (heavier) memory_debug_growth_mb: 50 # RSS delta since last tick → tracemalloc diff (if enabled) memory_debug_include_children: true # sum DataLoader worker RSS; pip install psutil When RSS keeps climbing, set memory_debug_interval: 1000–5000 and memory_debug_tracemalloc: true; logs will show file:line growth (often lists, caches, or Arrow/Python conversions). """ from __future__ import annotations import gc import logging import sys import tracemalloc logger = logging.getLogger(__name__) def current_rss_mb() -> float: """Resident set size in MiB for this process.""" try: with open("/proc/self/status", encoding="utf-8") as f: for line in f: if line.startswith("VmRSS:"): parts = line.split() return float(parts[1]) / 1024.0 # KiB → MiB except OSError: pass try: import psutil return psutil.Process().memory_info().rss / (1024.0 * 1024.0) except ImportError: import resource rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss if sys.platform == "darwin": return float(rss) / (1024.0 * 1024.0) return float(rss) / 1024.0 def current_rss_mb_with_children() -> float: """This process plus all descendants (e.g. DataLoader workers).""" total = current_rss_mb() try: import psutil p = psutil.Process() for c in p.children(recursive=True): try: total += c.memory_info().rss / (1024.0 * 1024.0) except (psutil.NoSuchProcess, psutil.AccessDenied): pass return total except ImportError: return total def cuda_mem_mb() -> tuple[float, float]: """Allocated and reserved CUDA memory (MiB) for current device.""" try: import torch if not torch.cuda.is_available(): return 0.0, 0.0 return ( torch.cuda.memory_allocated() / (1024.0 * 1024.0), torch.cuda.memory_reserved() / (1024.0 * 1024.0), ) except Exception: return 0.0, 0.0 class MemoryDebugSession: """RSS / CUDA logging; optional tracemalloc diff only when RSS jumps (cheap steady state).""" def __init__( self, rank: int, use_tracemalloc: bool = False, growth_alert_mb: float = 50.0, include_children: bool = True, ): self.rank = rank self.use_tracemalloc = use_tracemalloc self.growth_alert_mb = growth_alert_mb self.include_children = include_children self._prev_snapshot: tracemalloc.Snapshot | None = None self._prev_rss_mb: float | None = None self._started = False def _rss(self) -> float: if self.include_children: return current_rss_mb_with_children() return current_rss_mb() def start(self) -> None: if self._started: return self._started = True if self.use_tracemalloc: try: tracemalloc.start(25) except RuntimeError: logger.warning("[memdbg] tracemalloc was already started; continuing") self._prev_snapshot = tracemalloc.take_snapshot() self._prev_rss_mb = self._rss() def tick(self, step: int, force_tracemalloc_diff: bool = False) -> None: if not self._started: self.start() rss = self._rss() alloc_mb, res_mb = cuda_mem_mb() delta_rss = None if self._prev_rss_mb is None else rss - self._prev_rss_mb scope = "rss_tree_mib" if self.include_children else "rss_self_mib" if delta_rss is not None: logger.info( f"[memdbg rank={self.rank}] step={step} {scope}={rss:.1f} " f"delta_since_last_tick_mib={delta_rss:+.1f} " f"cuda_alloc_mib={alloc_mb:.1f} cuda_reserved_mib={res_mb:.1f}" ) else: logger.info( f"[memdbg rank={self.rank}] step={step} {scope}={rss:.1f} " f"cuda_alloc_mib={alloc_mb:.1f} cuda_reserved_mib={res_mb:.1f}" ) self._prev_rss_mb = rss if not self.use_tracemalloc: gc.collect() return growth = delta_rss is not None and delta_rss >= self.growth_alert_mb if growth or force_tracemalloc_diff: snap = tracemalloc.take_snapshot() if self._prev_snapshot is not None: stats = snap.compare_to(self._prev_snapshot, "lineno", cumulative=True) lines = [f"[memdbg rank={self.rank}] tracemalloc top growth (since last snapshot):"] for i, st in enumerate(stats[:25], 1): lines.append(f" {i}. {st}") logger.warning("\n".join(lines)) self._prev_snapshot = snap gc.collect()