"""memory_cleanup.py — Limpeza agressiva de memória estilo Xavante. Reaproveita os padrões de: - xavante_work/flexnet/advanced_memory_cleanup.py (AdvancedMemoryCleaner) - xavante_work/flexnet/oom_guard.py (OomGuard daemon) - xavante_work/xavante/utils/timing.py (TimeBudget) Implementa: 1. production_cleanup() — context manager com cleanup garantido 2. aggressive_cleanup() — gc.collect 3 gerações + torch.cuda.empty_cache 3. TimeBudget — orçamento de tempo por treino/época (streaming com timed steps) 4. get_rss_mb() — RSS do processo em MB """ from __future__ import annotations import gc import logging import os import threading import time from contextlib import contextmanager from dataclasses import dataclass, field from typing import Optional logger = logging.getLogger(__name__) def get_rss_mb() -> float: """Retorna o RSS (Resident Set Size) do processo atual em MB. Lê /proc/self/status (Linux). Fallback para psutil se disponível. """ try: with open("/proc/self/status", "r") as f: for line in f: if line.startswith("VmRSS:"): # VmRSS: 12345 kB parts = line.split() return float(parts[1]) / 1024.0 except (FileNotFoundError, IndexError, ValueError): pass # Fallback psutil try: import psutil return psutil.Process(os.getpid()).memory_info().rss / (1024 * 1024) except ImportError: return 0.0 def aggressive_cleanup(verbose: bool = False) -> dict: """Limpeza agressiva de memória (estilo Xavante). Sequência: 1. gc.collect gen 0, 1, 2 (3 gerações completas) 2. torch.cuda.empty_cache() (se CUDA disponível) 3. torch.cuda.synchronize() (se CUDA disponível) Args: verbose: logar memória antes/depois Returns: dict com rss_before_mb, rss_after_mb, freed_mb """ rss_before = get_rss_mb() # 3 gerações de gc gc.collect(0) gc.collect(1) gc.collect(2) # CUDA cleanup (no-op se CPU-only) try: import torch if torch.cuda.is_available(): torch.cuda.empty_cache() torch.cuda.synchronize() except Exception: pass rss_after = get_rss_mb() freed = rss_before - rss_after if verbose: logger.info( "aggressive_cleanup: RSS %.1f → %.1f MB (freed %.1f MB)", rss_before, rss_after, freed, ) return { "rss_before_mb": rss_before, "rss_after_mb": rss_after, "freed_mb": freed, } @contextmanager def production_cleanup(verbose: bool = False): """Context manager que garante cleanup agressivo ao sair (mesmo com exceção). Uso: with production_cleanup(verbose=True): # treino pesado aqui ... # cleanup automático ao sair do bloco """ try: yield finally: aggressive_cleanup(verbose=verbose) # --------------------------------------------------------------------------- # TimeBudget — orçamento de tempo para treino/época/passos # --------------------------------------------------------------------------- @dataclass class TimeBudget: """Orçamento de tempo para treino com timed steps. Permite: - max_total_s: tempo máximo total de treino - max_per_epoch_s: tempo máximo por época - is_over() / is_epoch_over(): verifica se orçamento estourou - remaining() / remaining_epoch(): tempo restante - reset_epoch(): reseta o contador de época (chamar no início de cada época) Uso: budget = TimeBudget(max_total_s=600, max_per_epoch_s=300) budget.reset_epoch() for epoch in range(2): for step in train_loop: if budget.is_over() or budget.is_epoch_over(): break ... budget.reset_epoch() """ max_total_s: float = 600.0 max_per_epoch_s: float = 300.0 _start: float = field(default_factory=time.time, repr=False) _epoch_start: float = field(default_factory=time.time, repr=False) def reset_epoch(self): """Reseta o contador de época (chamar no início de cada época).""" self._epoch_start = time.time() def reset_total(self): """Reseta o contador total (chamar no início do treino).""" self._start = time.time() self._epoch_start = time.time() def elapsed(self) -> float: return time.time() - self._start def elapsed_epoch(self) -> float: return time.time() - self._epoch_start def remaining(self) -> float: return max(0.0, self.max_total_s - self.elapsed()) def remaining_epoch(self) -> float: return max(0.0, self.max_per_epoch_s - self.elapsed_epoch()) def is_over(self) -> bool: return self.elapsed() >= self.max_total_s def is_epoch_over(self) -> bool: return self.elapsed_epoch() >= self.max_per_epoch_s def should_save_partial(self) -> bool: """True se faltam < 10% do tempo (salvar parcial).""" return self.remaining() < (self.max_total_s * 0.1) def summary(self) -> dict: return { "elapsed_s": self.elapsed(), "remaining_s": self.remaining(), "epoch_elapsed_s": self.elapsed_epoch(), "epoch_remaining_s": self.remaining_epoch(), "is_over": self.is_over(), "is_epoch_over": self.is_epoch_over(), } # --------------------------------------------------------------------------- # StepTimer — mede tempo por passo com warning se lento # --------------------------------------------------------------------------- @dataclass class StepTimer: """Mede tempo por passo e emite warnings se lento. Uso: timer = StepTimer(expected_s=0.5) for step in range(N): timer.start() # ... passo de treino ... timer.stop() # loga se > 2x expected """ expected_s: float = 0.5 _start: float = 0.0 _count: int = 0 _total_s: float = 0.0 _max_s: float = 0.0 def start(self): self._start = time.time() def stop(self, step_label: str = "") -> float: elapsed = time.time() - self._start self._count += 1 self._total_s += elapsed if elapsed > self._max_s: self._max_s = elapsed if elapsed > 2 * self.expected_s: logger.warning( "LENTO: step %s took %.2fs (expected ~%.2fs)", step_label or self._count, elapsed, self.expected_s, ) elif elapsed < 0.5 * self.expected_s: logger.debug("RAPIDO: step %s took %.2fs", step_label, elapsed) return elapsed def avg(self) -> float: return self._total_s / max(1, self._count) def summary(self) -> dict: return { "count": self._count, "total_s": self._total_s, "avg_s": self.avg(), "max_s": self._max_s, "expected_s": self.expected_s, } __all__ = [ "get_rss_mb", "aggressive_cleanup", "production_cleanup", "TimeBudget", "StepTimer", ]