V2: 12.97M params, HW optimizer, parallel BBPE, DPO, reasoning, inference, 10 bugs fixed
8594de8 verified Download src/bigru_t/utils/memory_cleanup.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 7.24 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/utils/memory_cleanup.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/utils/memory_cleanup.py
-
curl -L -o memory_cleanup.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/utils/memory_cleanup.py
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, | |
| } | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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", | |
| ] | |