BiGRU_T_version / src /bigru_t /utils /memory_cleanup.py
PowerMachine's picture
V2: 12.97M params, HW optimizer, parallel BBPE, DPO, reasoning, inference, 10 bugs fixed
8594de8 verified
Raw History Blame Contribute Delete
7.24 kB
"""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",
]