"""memory_optimizer.py — Otimização central de memória (CPU + GPU). Adaptado de xavante_work/xavante/utils/memory_optimizer.py, com: 1. gc.collect() explícito em checkpoints 2. torch.cuda.empty_cache() (se CUDA) 3. AMP (bfloat16/fp16) para reduzir VRAM 4. Gradient checkpointing 5. CPU offload de parâmetros congelados 6. Pinned memory para transferência async 7. set_per_process_memory_fraction (CUDA) Matemática: M_total = M_params + M_grads + M_activations + M_optimizer_state Com AMP: M_params *= 0.5, M_grads *= 0.5, M_activations *= 0.5 Com gradient checkpointing: M_activations *= 1/sqrt(L) """ from __future__ import annotations import gc import logging from contextlib import contextmanager from typing import Iterator import torch import torch.nn as nn logger = logging.getLogger(__name__) class MemoryOptimizer: """Centraliza otimização de memória para treino e inferência.""" def __init__(self, vram_fraction: float = 0.85, enable_amp: bool = True): self.vram_fraction = vram_fraction self.enable_amp = enable_amp self._peak_memory_mb: float = 0.0 def configure(self) -> None: """Configura PyTorch para uso otimizado de memória.""" try: torch.backends.cudnn.benchmark = True except Exception: pass try: torch.set_float32_matmul_precision("high") except Exception: pass if torch.cuda.is_available(): try: torch.cuda.set_per_process_memory_fraction(self.vram_fraction) logger.info("VRAM limitada a %.0f%%", self.vram_fraction * 100) except Exception as e: logger.warning("set_per_process_memory_fraction falhou: %s", e) @staticmethod def cleanup() -> None: """Limpa memória agressivamente (gc + empty_cache).""" gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() torch.cuda.synchronize() @staticmethod def get_memory_mb() -> dict: """Retorna uso de memória em MB.""" if torch.cuda.is_available(): return { "device": "cuda", "allocated_mb": torch.cuda.memory_allocated() / 1024**2, "cached_mb": torch.cuda.memory_reserved() / 1024**2, "max_allocated_mb": torch.cuda.max_memory_allocated() / 1024**2, } try: import psutil mem = psutil.virtual_memory() return { "device": "cpu", "total_mb": mem.total / 1024**2, "available_mb": mem.available / 1024**2, "used_mb": mem.used / 1024**2, "percent": mem.percent, } except ImportError: return {"device": "cpu", "info": "psutil not available"} @contextmanager def zero_grad_context(self, model: nn.Module) -> Iterator[None]: """Context manager que limpa gradientes ao sair.""" try: yield finally: model.zero_grad(set_to_none=True) self.cleanup() @staticmethod def enable_gradient_checkpointing(model: nn.Module) -> None: """Tenta ativar gradient checkpointing.""" if hasattr(model, "gradient_checkpointing_enable"): try: model.gradient_checkpointing_enable() logger.info("Gradient checkpointing ativado") return except Exception as e: logger.warning("gradient_checkpointing_enable falhou: %s", e) logger.info("Modelo não suporta gradient checkpointing nativo") @staticmethod def cpu_offload_constrained(model: nn.Module) -> int: """Move parâmetros sem requires_grad para CPU.""" n_offloaded = 0 for p in model.parameters(): if not p.requires_grad and p.device.type != "cpu": p.data = p.data.cpu() n_offloaded += p.numel() if n_offloaded > 0: logger.info("Offloaded %d params para CPU", n_offloaded) if torch.cuda.is_available(): torch.cuda.empty_cache() return n_offloaded def amp_context(self, device_type: str = "cpu"): """Context manager para mixed precision.""" if not self.enable_amp or device_type == "cpu": from contextlib import nullcontext return nullcontext() try: return torch.amp.autocast(device_type=device_type, dtype=torch.bfloat16) except Exception: from contextlib import nullcontext return nullcontext() @staticmethod def report_peak() -> float: if torch.cuda.is_available(): return torch.cuda.max_memory_allocated() / 1024**2 return 0.0 __all__ = ["MemoryOptimizer"]