Download cnn_bigru/utils/memory_optimizer.py from PowerMachine/CNN-BiGRU: direct link, hf CLI and curl.
- Browser
- Download file 4.91 kB
-
https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/utils/memory_optimizer.py
- Command line
-
hf download hf://PowerMachine/CNN-BiGRU/cnn_bigru/utils/memory_optimizer.py
-
curl -L -o memory_optimizer.py https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/utils/memory_optimizer.py
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) | |
| def cleanup() -> None: | |
| """Limpa memória agressivamente (gc + empty_cache).""" | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| torch.cuda.synchronize() | |
| 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"} | |
| 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() | |
| 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") | |
| 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() | |
| def report_peak() -> float: | |
| if torch.cuda.is_available(): | |
| return torch.cuda.max_memory_allocated() / 1024**2 | |
| return 0.0 | |
| __all__ = ["MemoryOptimizer"] | |