| """ |
| 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 |
| 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() |
|
|