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