CNN-BiGRU / cnn_bigru /utils /memory_optimizer.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw
History Blame Contribute Delete
4.91 kB
"""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"]