| """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"] |
|
|