File size: 4,909 Bytes
49b8205 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 | """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"]
|