MacroCast / tempopfn /src /utils /memory_debug.py
shubhranshu's picture
Upload MacroCast (26 vintage checkpoints)
294e473 verified
Raw
History Blame Contribute Delete
5.18 kB
"""
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()