mztrain 1.3.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mztrain/__init__.py +176 -0
- mztrain/adaptive_optimizer.py +420 -0
- mztrain/checkpoint.py +263 -0
- mztrain/config.py +585 -0
- mztrain/data/__init__.py +17 -0
- mztrain/data/dataset.py +576 -0
- mztrain/data/tokenizer.py +547 -0
- mztrain/elastic_rank.py +968 -0
- mztrain/elastic_shape.py +416 -0
- mztrain/engine.py +1639 -0
- mztrain/gradient.py +383 -0
- mztrain/layers.py +852 -0
- mztrain/models/__init__.py +28 -0
- mztrain/models/zcodebert.py +835 -0
- mztrain/optimizer.py +295 -0
- mztrain/precision.py +529 -0
- mztrain/projector.py +462 -0
- mztrain/refactorize.py +185 -0
- mztrain/scheduler.py +376 -0
- mztrain/shape_ops.py +269 -0
- mztrain/utils.py +216 -0
- mztrain/vram_governor.py +243 -0
- mztrain-1.3.0.dist-info/METADATA +410 -0
- mztrain-1.3.0.dist-info/RECORD +29 -0
- mztrain-1.3.0.dist-info/WHEEL +5 -0
- mztrain-1.3.0.dist-info/licenses/AUTHORSHIP.md +137 -0
- mztrain-1.3.0.dist-info/licenses/LICENSE +21 -0
- mztrain-1.3.0.dist-info/licenses/NOTICE.md +63 -0
- mztrain-1.3.0.dist-info/top_level.txt +1 -0
mztrain/__init__.py
ADDED
|
@@ -0,0 +1,176 @@
|
|
|
1
|
+
"""
|
|
2
|
+
MZTrain - Motor de Entrenamiento en Espacio Comprimido v1.0
|
|
3
|
+
|
|
4
|
+
Tecnologia que permite entrenar modelos de IA sin miles de GPUs.
|
|
5
|
+
En vez de entrenar tensores completos (W: m x n), entrena sus factores
|
|
6
|
+
descompuestos (U: m x r, S: r, V: r x n) donde r << min(m,n).
|
|
7
|
+
|
|
8
|
+
Principio: Si MNEME puede comprimir un modelo entrenado 93%+, entonces
|
|
9
|
+
podemos ENTRENAR directamente en ese espacio comprimido.
|
|
10
|
+
|
|
11
|
+
Componentes:
|
|
12
|
+
- ZFactorizedLinear: Entrena factores SVD directamente
|
|
13
|
+
- ZFactorizedAttention: Multi-head attention factorizado
|
|
14
|
+
- ZFactorizedTransformerBlock: Bloque transformer completo factorizado
|
|
15
|
+
- ZCompressedAdam: Optimizer con estados comprimidos INT8
|
|
16
|
+
- ZGradientCompressor: Compresion de gradientes con error feedback
|
|
17
|
+
- ZRankScheduler: Crecimiento progresivo de rango
|
|
18
|
+
- ZActivationCheckpoint: Checkpointing comprimido de activaciones
|
|
19
|
+
- ZTrainEngine: Orquestador principal del entrenamiento
|
|
20
|
+
|
|
21
|
+
Ahorro de memoria vs entrenamiento tradicional:
|
|
22
|
+
- Pesos: ~20-30% del original
|
|
23
|
+
- Gradientes: ~20-30% del original
|
|
24
|
+
- Optimizer: ~20-30% del original
|
|
25
|
+
- Activaciones: ~40-60% del original
|
|
26
|
+
- TOTAL: Un modelo de 1B params que necesita 32GB VRAM
|
|
27
|
+
ahora necesita ~6-10GB VRAM
|
|
28
|
+
|
|
29
|
+
Autores: MSC Star Team (Esraderey y Raul Cruz Acosta)
|
|
30
|
+
Basado en: MNEME Motor de Memoria Neural Morfica v2.0
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
from .config import (
|
|
34
|
+
ZTrainConfig,
|
|
35
|
+
RankSchedule,
|
|
36
|
+
GradientCompression,
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
from .layers import (
|
|
40
|
+
ZFactorizedLinear,
|
|
41
|
+
ZFactorizedAttention,
|
|
42
|
+
ZFactorizedTransformerBlock,
|
|
43
|
+
ZSparseFactorizedLinear,
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
from .engine import ZTrainEngine
|
|
47
|
+
|
|
48
|
+
from .optimizer import ZCompressedAdam
|
|
49
|
+
|
|
50
|
+
from .gradient import ZGradientCompressor
|
|
51
|
+
|
|
52
|
+
from .projector import ZGaLoreProjector, ZGaLoreOptimizer
|
|
53
|
+
|
|
54
|
+
from .adaptive_optimizer import ZAdaptiveOptimizer
|
|
55
|
+
|
|
56
|
+
from .precision import ZMultiPrecisionManager
|
|
57
|
+
|
|
58
|
+
from .scheduler import ZRankScheduler, ZSpectralRankScheduler
|
|
59
|
+
|
|
60
|
+
from .checkpoint import (
|
|
61
|
+
ZActivationCheckpoint,
|
|
62
|
+
z_checkpoint,
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
from .refactorize import refactorize_model
|
|
66
|
+
|
|
67
|
+
from .elastic_rank import (
|
|
68
|
+
ElasticRankController,
|
|
69
|
+
SleepingDirection,
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
from .vram_governor import ZVRAMGovernor
|
|
73
|
+
|
|
74
|
+
from .utils import (
|
|
75
|
+
factorize_existing_model,
|
|
76
|
+
estimate_memory_savings,
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
# Sub-packages
|
|
80
|
+
from . import data
|
|
81
|
+
from . import models
|
|
82
|
+
from .models.zcodebert import (
|
|
83
|
+
ZCodeBERTConfig,
|
|
84
|
+
ZCodeBERT,
|
|
85
|
+
ZCodeBERTForMLM,
|
|
86
|
+
ZCodeBERTForCausalLM,
|
|
87
|
+
ZCodeBERTForSequenceClassification,
|
|
88
|
+
)
|
|
89
|
+
from .data.tokenizer import CodeTokenizer
|
|
90
|
+
from .data.dataset import SyntheticCodeGenerator, CodeDataset, MLMDataset, CausalCodeDataset
|
|
91
|
+
from .shape_ops import (
|
|
92
|
+
dense_to_factorized,
|
|
93
|
+
pad_adam_entry,
|
|
94
|
+
pad_state_tensor,
|
|
95
|
+
scale_output_rows,
|
|
96
|
+
widen_embedding,
|
|
97
|
+
widen_factorized_linear,
|
|
98
|
+
widen_layernorm,
|
|
99
|
+
zero_block_outputs,
|
|
100
|
+
)
|
|
101
|
+
from .elastic_shape import (
|
|
102
|
+
GrowthEvent,
|
|
103
|
+
LrWarmup,
|
|
104
|
+
ShapeSchedule,
|
|
105
|
+
apply_event,
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
__version__ = "1.3.0"
|
|
109
|
+
__author__ = "MSC Star Team (Esraderey y Raul Cruz Acosta)"
|
|
110
|
+
__email__ = "msc.framework@gmail.com"
|
|
111
|
+
|
|
112
|
+
__all__ = [
|
|
113
|
+
# ElasticShape v1 (crecimiento de forma; claim T8 en docs/evidencia/)
|
|
114
|
+
"dense_to_factorized",
|
|
115
|
+
"pad_adam_entry",
|
|
116
|
+
"pad_state_tensor",
|
|
117
|
+
"scale_output_rows",
|
|
118
|
+
"widen_embedding",
|
|
119
|
+
"widen_factorized_linear",
|
|
120
|
+
"widen_layernorm",
|
|
121
|
+
"zero_block_outputs",
|
|
122
|
+
"GrowthEvent",
|
|
123
|
+
"LrWarmup",
|
|
124
|
+
"ShapeSchedule",
|
|
125
|
+
"apply_event",
|
|
126
|
+
# Config
|
|
127
|
+
"ZTrainConfig",
|
|
128
|
+
"RankSchedule",
|
|
129
|
+
"GradientCompression",
|
|
130
|
+
# Layers
|
|
131
|
+
"ZFactorizedLinear",
|
|
132
|
+
"ZFactorizedAttention",
|
|
133
|
+
"ZFactorizedTransformerBlock",
|
|
134
|
+
"ZSparseFactorizedLinear",
|
|
135
|
+
# Engine
|
|
136
|
+
"ZTrainEngine",
|
|
137
|
+
# Optimizer
|
|
138
|
+
"ZCompressedAdam",
|
|
139
|
+
# Gradient
|
|
140
|
+
"ZGradientCompressor",
|
|
141
|
+
# Projector (GaLore)
|
|
142
|
+
"ZGaLoreProjector",
|
|
143
|
+
"ZGaLoreOptimizer",
|
|
144
|
+
# Adaptive Optimizer (APOLLO, T7)
|
|
145
|
+
"ZAdaptiveOptimizer",
|
|
146
|
+
# Precision Manager (T8)
|
|
147
|
+
"ZMultiPrecisionManager",
|
|
148
|
+
# Scheduler
|
|
149
|
+
"ZRankScheduler",
|
|
150
|
+
"ZSpectralRankScheduler",
|
|
151
|
+
# Checkpointing
|
|
152
|
+
"ZActivationCheckpoint",
|
|
153
|
+
"z_checkpoint",
|
|
154
|
+
# Refactorize
|
|
155
|
+
"refactorize_model",
|
|
156
|
+
# ElasticRank (rango bidireccional)
|
|
157
|
+
"ElasticRankController",
|
|
158
|
+
"SleepingDirection",
|
|
159
|
+
# VRAM Governor (MVP: gating de growth + OOM-retry)
|
|
160
|
+
"ZVRAMGovernor",
|
|
161
|
+
# Utilities
|
|
162
|
+
"factorize_existing_model",
|
|
163
|
+
"estimate_memory_savings",
|
|
164
|
+
# Models
|
|
165
|
+
"ZCodeBERTConfig",
|
|
166
|
+
"ZCodeBERT",
|
|
167
|
+
"ZCodeBERTForMLM",
|
|
168
|
+
"ZCodeBERTForCausalLM",
|
|
169
|
+
"ZCodeBERTForSequenceClassification",
|
|
170
|
+
# Data
|
|
171
|
+
"CodeTokenizer",
|
|
172
|
+
"SyntheticCodeGenerator",
|
|
173
|
+
"CodeDataset",
|
|
174
|
+
"MLMDataset",
|
|
175
|
+
"CausalCodeDataset",
|
|
176
|
+
]
|
|
@@ -0,0 +1,420 @@
|
|
|
1
|
+
"""
|
|
2
|
+
MZTrain - APOLLO-style Adaptive Optimizer (T7).
|
|
3
|
+
|
|
4
|
+
Combina lo mejor de tres papers:
|
|
5
|
+
- GaLore (ICML 2024): Proyeccion low-rank del primer momento m
|
|
6
|
+
- Adam-mini (ICLR 2025): Segundo momento v escalar por bloque
|
|
7
|
+
- APOLLO (MLSys 2025 Outstanding Paper): Combinacion con error compensation
|
|
8
|
+
|
|
9
|
+
Resultado: memoria ~= SGD, convergencia ~= AdamW.
|
|
10
|
+
|
|
11
|
+
El cuello de botella de memoria en modelos >1B params NO son los pesos
|
|
12
|
+
(ya factorizados), son los estados del optimizer. Adam guarda m y v del
|
|
13
|
+
mismo tamano que los parametros. ZAdaptiveOptimizer reduce esto ~96.8%:
|
|
14
|
+
|
|
15
|
+
Adam estandar (1B): m=2.15GB + v=2.15GB = 4.3GB
|
|
16
|
+
ZAdaptiveOptimizer: m_lr=67MB + v_scalar=2MB + P=67MB = ~136MB
|
|
17
|
+
|
|
18
|
+
Para parametros pequenos (<2D o <4096 elementos), usa Adam estandar
|
|
19
|
+
completo para preservar convergencia.
|
|
20
|
+
|
|
21
|
+
Referencias:
|
|
22
|
+
- APOLLO: SGD-level memory, AdamW-level performance (MLSys 2025)
|
|
23
|
+
- Adam-mini: Zhang et al., ICLR 2025 (arXiv:2406.16793)
|
|
24
|
+
- GaLore: Zhao et al., ICML 2024 Oral (arXiv:2403.03507)
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
import math
|
|
28
|
+
import logging
|
|
29
|
+
from typing import Optional
|
|
30
|
+
|
|
31
|
+
logger = logging.getLogger("mztrain")
|
|
32
|
+
|
|
33
|
+
try:
|
|
34
|
+
import torch
|
|
35
|
+
HAS_TORCH = True
|
|
36
|
+
except ImportError:
|
|
37
|
+
HAS_TORCH = False
|
|
38
|
+
|
|
39
|
+
if not HAS_TORCH:
|
|
40
|
+
raise ImportError(
|
|
41
|
+
"PyTorch es requerido para mztrain.adaptive_optimizer. "
|
|
42
|
+
"Instalar con: pip install torch>=2.0.0"
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
from .projector import ZGaLoreProjector
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class ZAdaptiveOptimizer(torch.optim.Optimizer):
|
|
49
|
+
"""Adam optimizer APOLLO-style: m low-rank + v escalar por bloque.
|
|
50
|
+
|
|
51
|
+
Para parametros 2D grandes (>= min_dim_to_project elementos):
|
|
52
|
+
- m (primer momento): proyectado a subespacio de rango r via GaLore
|
|
53
|
+
- v (segundo momento): UN SOLO ESCALAR por bloque de block_size elementos
|
|
54
|
+
- Error compensation: residuo de la proyeccion se acumula
|
|
55
|
+
|
|
56
|
+
Para parametros pequenos o 1D (biases, layer norms):
|
|
57
|
+
- Adam estandar completo (m y v per-element)
|
|
58
|
+
|
|
59
|
+
La diferencia clave vs ZGaLoreOptimizer es que v es escalar por bloque
|
|
60
|
+
en vez de per-element en espacio proyectado. Esto ahorra >96% en v.
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
params: Parametros del modelo.
|
|
64
|
+
lr: Learning rate.
|
|
65
|
+
betas: Coeficientes de momentos (beta1, beta2).
|
|
66
|
+
eps: Epsilon para estabilidad numerica.
|
|
67
|
+
weight_decay: Weight decay (decoupled, estilo AdamW).
|
|
68
|
+
rank: Rango del subespacio de proyeccion para m.
|
|
69
|
+
block_size: Tamano de bloque para v escalar (1 valor v cada N elementos).
|
|
70
|
+
projection_update_freq: Cada cuantos steps actualizar subespacio via SVD.
|
|
71
|
+
min_dim_to_project: Dimension minima para aplicar proyeccion.
|
|
72
|
+
compress_states: Comprimir m low-rank con INT8 block-wise.
|
|
73
|
+
compression_interval: Re-comprimir cada N steps.
|
|
74
|
+
|
|
75
|
+
Example:
|
|
76
|
+
>>> optimizer = ZAdaptiveOptimizer(model.parameters(), lr=3e-4, rank=128)
|
|
77
|
+
>>> optimizer.step()
|
|
78
|
+
>>> print(optimizer.get_memory_stats())
|
|
79
|
+
"""
|
|
80
|
+
|
|
81
|
+
_BLOCK_SIZE_COMPRESS = 2048 # Block size para INT8 compression de estados
|
|
82
|
+
|
|
83
|
+
def __init__(
|
|
84
|
+
self,
|
|
85
|
+
params,
|
|
86
|
+
lr: float = 3e-4,
|
|
87
|
+
betas: tuple = (0.9, 0.999),
|
|
88
|
+
eps: float = 1e-8,
|
|
89
|
+
weight_decay: float = 1e-4,
|
|
90
|
+
rank: int = 128,
|
|
91
|
+
block_size: int = 1024,
|
|
92
|
+
projection_update_freq: int = 200,
|
|
93
|
+
min_dim_to_project: int = 128,
|
|
94
|
+
compress_states: bool = False,
|
|
95
|
+
compression_interval: int = 10,
|
|
96
|
+
):
|
|
97
|
+
defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay)
|
|
98
|
+
super().__init__(params, defaults)
|
|
99
|
+
self.rank = rank
|
|
100
|
+
self.block_size = block_size
|
|
101
|
+
self.projection_update_freq = projection_update_freq
|
|
102
|
+
self.min_dim_to_project = min_dim_to_project
|
|
103
|
+
self.compress_states = compress_states
|
|
104
|
+
self.compression_interval = compression_interval
|
|
105
|
+
self._step_count = 0
|
|
106
|
+
|
|
107
|
+
def _should_project(self, param: torch.Tensor) -> bool:
|
|
108
|
+
"""Determinar si un parametro debe usar proyeccion APOLLO."""
|
|
109
|
+
if param.dim() < 2:
|
|
110
|
+
return False
|
|
111
|
+
if min(param.shape[0], param.shape[1]) < self.min_dim_to_project:
|
|
112
|
+
return False
|
|
113
|
+
if param.numel() < 4096:
|
|
114
|
+
return False
|
|
115
|
+
return True
|
|
116
|
+
|
|
117
|
+
def _compute_n_blocks(self, numel: int) -> int:
|
|
118
|
+
"""Calcular numero de bloques para v escalar."""
|
|
119
|
+
return max(1, (numel + self.block_size - 1) // self.block_size)
|
|
120
|
+
|
|
121
|
+
def _compute_block_variance(self, grad_low: torch.Tensor) -> torch.Tensor:
|
|
122
|
+
"""Calcular varianza (g^2) promedio por bloque del gradiente.
|
|
123
|
+
|
|
124
|
+
Divide el gradiente aplanado en bloques de block_size y calcula
|
|
125
|
+
la media de g^2 en cada bloque. Esto es la estimacion de varianza
|
|
126
|
+
escalar estilo Adam-mini.
|
|
127
|
+
|
|
128
|
+
Args:
|
|
129
|
+
grad_low: Gradiente en espacio low-rank (cualquier shape).
|
|
130
|
+
|
|
131
|
+
Returns:
|
|
132
|
+
Tensor de shape (n_blocks,) con la varianza promedio por bloque.
|
|
133
|
+
"""
|
|
134
|
+
flat = grad_low.reshape(-1)
|
|
135
|
+
numel = flat.numel()
|
|
136
|
+
n_blocks = self._compute_n_blocks(numel)
|
|
137
|
+
|
|
138
|
+
# Pad si no es multiplo exacto
|
|
139
|
+
if numel % self.block_size != 0:
|
|
140
|
+
padded = torch.zeros(n_blocks * self.block_size,
|
|
141
|
+
device=flat.device, dtype=flat.dtype)
|
|
142
|
+
padded[:numel] = flat
|
|
143
|
+
else:
|
|
144
|
+
padded = flat
|
|
145
|
+
|
|
146
|
+
# Reshape a (n_blocks, block_size) y calcular mean(g^2) por bloque
|
|
147
|
+
blocks = padded.reshape(n_blocks, -1)
|
|
148
|
+
block_var = blocks.pow(2).mean(dim=1) # (n_blocks,)
|
|
149
|
+
|
|
150
|
+
return block_var
|
|
151
|
+
|
|
152
|
+
def _expand_block_scalar(
|
|
153
|
+
self, v_scalar: torch.Tensor, target_shape: tuple
|
|
154
|
+
) -> torch.Tensor:
|
|
155
|
+
"""Expandir v escalar por bloque a la shape del gradiente.
|
|
156
|
+
|
|
157
|
+
Cada valor escalar se repite block_size veces para cubrir
|
|
158
|
+
el gradiente completo.
|
|
159
|
+
|
|
160
|
+
Args:
|
|
161
|
+
v_scalar: Tensor (n_blocks,) con varianzas escalares.
|
|
162
|
+
target_shape: Shape objetivo del gradiente.
|
|
163
|
+
|
|
164
|
+
Returns:
|
|
165
|
+
Tensor de target_shape con valores escalares expandidos.
|
|
166
|
+
"""
|
|
167
|
+
numel = 1
|
|
168
|
+
for s in target_shape:
|
|
169
|
+
numel *= s
|
|
170
|
+
|
|
171
|
+
# Expandir: cada escalar -> block_size elementos
|
|
172
|
+
expanded = v_scalar.unsqueeze(1).expand(-1, self.block_size).reshape(-1)
|
|
173
|
+
expanded = expanded[:numel]
|
|
174
|
+
|
|
175
|
+
return expanded.reshape(target_shape)
|
|
176
|
+
|
|
177
|
+
# === INT8 Compression (reutilizado de ZGaLoreOptimizer) ===
|
|
178
|
+
|
|
179
|
+
def _compress_state(self, tensor: torch.Tensor) -> tuple:
|
|
180
|
+
"""Comprimir estado a INT8 block-wise."""
|
|
181
|
+
flat = tensor.reshape(-1)
|
|
182
|
+
n = flat.numel()
|
|
183
|
+
bs = self._BLOCK_SIZE_COMPRESS
|
|
184
|
+
n_blocks = (n + bs - 1) // bs
|
|
185
|
+
|
|
186
|
+
if n % bs != 0:
|
|
187
|
+
padded = torch.zeros(n_blocks * bs, device=flat.device, dtype=flat.dtype)
|
|
188
|
+
padded[:n] = flat
|
|
189
|
+
else:
|
|
190
|
+
padded = flat
|
|
191
|
+
|
|
192
|
+
blocks = padded.reshape(n_blocks, bs)
|
|
193
|
+
scales = blocks.abs().amax(dim=1)
|
|
194
|
+
safe_scales = scales.clamp(min=1e-12)
|
|
195
|
+
# Use full INT8 range: map [-abs_max, abs_max] -> [-127, 127].
|
|
196
|
+
# Without the *127 factor, values collapse to {-1, 0, 1} and exp_avg
|
|
197
|
+
# is destroyed on every compression cycle.
|
|
198
|
+
quantized = (
|
|
199
|
+
(blocks / safe_scales.unsqueeze(1) * 127.0)
|
|
200
|
+
.round().clamp(-127, 127).to(torch.int8)
|
|
201
|
+
)
|
|
202
|
+
|
|
203
|
+
return quantized, scales, tensor.shape, n
|
|
204
|
+
|
|
205
|
+
def _decompress_state(self, compressed: tuple, dtype: torch.dtype) -> torch.Tensor:
|
|
206
|
+
"""Descomprimir estado desde INT8 block-wise."""
|
|
207
|
+
quantized, scales, shape, orig_numel = compressed
|
|
208
|
+
# Inverse of compress: divide by 127 to undo the INT8 range scaling.
|
|
209
|
+
blocks = quantized.to(dtype) * (scales.to(dtype) / 127.0).unsqueeze(1)
|
|
210
|
+
flat = blocks.reshape(-1)[:orig_numel]
|
|
211
|
+
return flat.reshape(shape)
|
|
212
|
+
|
|
213
|
+
@torch.no_grad()
|
|
214
|
+
def step(self, closure=None):
|
|
215
|
+
"""Paso de optimizacion APOLLO-style.
|
|
216
|
+
|
|
217
|
+
Parametros grandes (2D, >4096 elementos):
|
|
218
|
+
1. Proyectar gradiente: g_low = P^T @ grad (r x n)
|
|
219
|
+
2. Actualizar m low-rank: m = beta1*m + (1-beta1)*g_low
|
|
220
|
+
3. Actualizar v escalar: v_block = beta2*v + (1-beta2)*mean(g_low^2, per_block)
|
|
221
|
+
4. Calcular update: update = m / (sqrt(v_expanded) + eps)
|
|
222
|
+
5. Reconstruir: delta_W = P @ update
|
|
223
|
+
|
|
224
|
+
Parametros pequenos (1D, biases, norms):
|
|
225
|
+
Adam estandar completo.
|
|
226
|
+
|
|
227
|
+
Args:
|
|
228
|
+
closure: Closure para reevaluar loss (opcional).
|
|
229
|
+
|
|
230
|
+
Returns:
|
|
231
|
+
Loss si closure fue proporcionado, None de lo contrario.
|
|
232
|
+
"""
|
|
233
|
+
loss = None
|
|
234
|
+
if closure is not None:
|
|
235
|
+
with torch.enable_grad():
|
|
236
|
+
loss = closure()
|
|
237
|
+
|
|
238
|
+
self._step_count += 1
|
|
239
|
+
|
|
240
|
+
for group in self.param_groups:
|
|
241
|
+
beta1, beta2 = group['betas']
|
|
242
|
+
lr = group['lr']
|
|
243
|
+
eps = group['eps']
|
|
244
|
+
wd = group['weight_decay']
|
|
245
|
+
|
|
246
|
+
for p in group['params']:
|
|
247
|
+
if p.grad is None:
|
|
248
|
+
continue
|
|
249
|
+
|
|
250
|
+
grad = p.grad
|
|
251
|
+
state = self.state[p]
|
|
252
|
+
|
|
253
|
+
# === LAZY INIT ===
|
|
254
|
+
if len(state) == 0:
|
|
255
|
+
state['step'] = 0
|
|
256
|
+
|
|
257
|
+
if self._should_project(p):
|
|
258
|
+
m, n = p.shape[0], p.shape[1]
|
|
259
|
+
proj = ZGaLoreProjector(
|
|
260
|
+
m, n, self.rank, self.projection_update_freq,
|
|
261
|
+
)
|
|
262
|
+
state['projector'] = proj
|
|
263
|
+
lr_shape = proj.low_rank_shape
|
|
264
|
+
|
|
265
|
+
# m: low-rank (r x n) o (m x r)
|
|
266
|
+
state['exp_avg'] = torch.zeros(
|
|
267
|
+
lr_shape, device=p.device, dtype=p.dtype
|
|
268
|
+
)
|
|
269
|
+
|
|
270
|
+
# v: escalar por bloque — UN valor cada block_size elementos
|
|
271
|
+
lr_numel = lr_shape[0] * lr_shape[1]
|
|
272
|
+
n_blocks = self._compute_n_blocks(lr_numel)
|
|
273
|
+
state['exp_avg_sq'] = torch.zeros(
|
|
274
|
+
n_blocks, device=p.device, dtype=p.dtype
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
state['is_projected'] = True
|
|
278
|
+
else:
|
|
279
|
+
# Adam estandar para params pequenos
|
|
280
|
+
state['exp_avg'] = torch.zeros_like(p.data)
|
|
281
|
+
state['exp_avg_sq'] = torch.zeros_like(p.data)
|
|
282
|
+
state['is_projected'] = False
|
|
283
|
+
|
|
284
|
+
state['compressed'] = False
|
|
285
|
+
|
|
286
|
+
state['step'] += 1
|
|
287
|
+
t = state['step']
|
|
288
|
+
|
|
289
|
+
# Descomprimir si necesario
|
|
290
|
+
if state['compressed'] and self.compress_states:
|
|
291
|
+
m_val = self._decompress_state(state['exp_avg'], p.dtype)
|
|
292
|
+
# v escalar NO se comprime (ya es tiny)
|
|
293
|
+
v_val = state['exp_avg_sq']
|
|
294
|
+
if isinstance(v_val, tuple):
|
|
295
|
+
v_val = self._decompress_state(v_val, p.dtype)
|
|
296
|
+
else:
|
|
297
|
+
m_val = state['exp_avg']
|
|
298
|
+
v_val = state['exp_avg_sq']
|
|
299
|
+
|
|
300
|
+
# Weight decay (decoupled, estilo AdamW)
|
|
301
|
+
if wd != 0:
|
|
302
|
+
p.data.mul_(1 - lr * wd)
|
|
303
|
+
|
|
304
|
+
if state['is_projected']:
|
|
305
|
+
# === APOLLO PROJECTED PATH ===
|
|
306
|
+
proj = state['projector']
|
|
307
|
+
grad_2d = grad.reshape(p.shape[0], -1)
|
|
308
|
+
|
|
309
|
+
# 1. Proyectar gradiente al subespacio low-rank
|
|
310
|
+
g_low = proj.project(grad_2d)
|
|
311
|
+
|
|
312
|
+
# 2. Actualizar m (low-rank, per-element en espacio proyectado)
|
|
313
|
+
m_val.mul_(beta1).add_(g_low, alpha=1 - beta1)
|
|
314
|
+
|
|
315
|
+
# 3. Actualizar v (ESCALAR por bloque — clave de APOLLO)
|
|
316
|
+
block_var = self._compute_block_variance(g_low)
|
|
317
|
+
v_val.mul_(beta2).add_(block_var, alpha=1 - beta2)
|
|
318
|
+
|
|
319
|
+
# 4. Bias correction
|
|
320
|
+
bc1 = 1 - beta1 ** t
|
|
321
|
+
bc2 = 1 - beta2 ** t
|
|
322
|
+
step_size = lr / bc1
|
|
323
|
+
|
|
324
|
+
# 5. Expandir v escalar a shape de m para division
|
|
325
|
+
v_expanded = self._expand_block_scalar(v_val, m_val.shape)
|
|
326
|
+
denom = (v_expanded.sqrt() / math.sqrt(bc2)).add_(eps)
|
|
327
|
+
|
|
328
|
+
# 6. Update en espacio low-rank
|
|
329
|
+
update_low = m_val / denom
|
|
330
|
+
|
|
331
|
+
# 7. Reconstruir en espacio completo y aplicar
|
|
332
|
+
update_full = proj.project_back(update_low * (-step_size))
|
|
333
|
+
p.data.add_(update_full.reshape(p.shape))
|
|
334
|
+
|
|
335
|
+
else:
|
|
336
|
+
# === STANDARD ADAM PATH (params pequenos) ===
|
|
337
|
+
m_val.mul_(beta1).add_(grad, alpha=1 - beta1)
|
|
338
|
+
v_val.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
|
|
339
|
+
|
|
340
|
+
bc1 = 1 - beta1 ** t
|
|
341
|
+
bc2 = 1 - beta2 ** t
|
|
342
|
+
step_size = lr / bc1
|
|
343
|
+
denom = (v_val.sqrt() / math.sqrt(bc2)).add_(eps)
|
|
344
|
+
|
|
345
|
+
p.data.addcdiv_(m_val, denom, value=-step_size)
|
|
346
|
+
|
|
347
|
+
# Re-comprimir periodicamente (solo m, v escalar es tiny)
|
|
348
|
+
should_compress = (
|
|
349
|
+
self.compress_states
|
|
350
|
+
and self._step_count % self.compression_interval == 0
|
|
351
|
+
and state['is_projected']
|
|
352
|
+
)
|
|
353
|
+
if should_compress:
|
|
354
|
+
state['exp_avg'] = self._compress_state(m_val)
|
|
355
|
+
state['exp_avg_sq'] = v_val # v escalar ya es minusculo
|
|
356
|
+
state['compressed'] = True
|
|
357
|
+
else:
|
|
358
|
+
state['exp_avg'] = m_val
|
|
359
|
+
state['exp_avg_sq'] = v_val
|
|
360
|
+
state['compressed'] = False
|
|
361
|
+
|
|
362
|
+
return loss
|
|
363
|
+
|
|
364
|
+
def get_memory_stats(self) -> dict:
|
|
365
|
+
"""Estadisticas de memoria del optimizer.
|
|
366
|
+
|
|
367
|
+
Returns:
|
|
368
|
+
Dict con metricas de memoria: total_params, projected/standard counts,
|
|
369
|
+
state_memory_mb, full_memory_mb, memory_saved_pct, rank, block_size.
|
|
370
|
+
"""
|
|
371
|
+
total_state_bytes = 0
|
|
372
|
+
total_full_bytes = 0
|
|
373
|
+
projected_count = 0
|
|
374
|
+
standard_count = 0
|
|
375
|
+
|
|
376
|
+
for group in self.param_groups:
|
|
377
|
+
for p in group['params']:
|
|
378
|
+
state = self.state.get(p, {})
|
|
379
|
+
if len(state) == 0:
|
|
380
|
+
continue
|
|
381
|
+
|
|
382
|
+
param_bytes = p.numel() * p.element_size()
|
|
383
|
+
total_full_bytes += param_bytes * 2 # m + v full size
|
|
384
|
+
|
|
385
|
+
if state.get('is_projected', False):
|
|
386
|
+
projected_count += 1
|
|
387
|
+
proj = state['projector']
|
|
388
|
+
lr_numel = proj.low_rank_shape[0] * proj.low_rank_shape[1]
|
|
389
|
+
n_blocks = self._compute_n_blocks(lr_numel)
|
|
390
|
+
|
|
391
|
+
if state.get('compressed', False):
|
|
392
|
+
# m comprimido a INT8
|
|
393
|
+
m_bytes = lr_numel * 1 # INT8
|
|
394
|
+
else:
|
|
395
|
+
m_bytes = lr_numel * p.element_size()
|
|
396
|
+
|
|
397
|
+
# v escalar: n_blocks floats (minusculo)
|
|
398
|
+
v_bytes = n_blocks * p.element_size()
|
|
399
|
+
|
|
400
|
+
# Projector storage (P matrix)
|
|
401
|
+
proj_bytes = proj.m * proj.rank * p.element_size()
|
|
402
|
+
|
|
403
|
+
total_state_bytes += m_bytes + v_bytes + proj_bytes
|
|
404
|
+
else:
|
|
405
|
+
standard_count += 1
|
|
406
|
+
total_state_bytes += param_bytes * 2 # m + v full
|
|
407
|
+
|
|
408
|
+
return {
|
|
409
|
+
"total_params": projected_count + standard_count,
|
|
410
|
+
"projected_params": projected_count,
|
|
411
|
+
"standard_params": standard_count,
|
|
412
|
+
"state_memory_mb": total_state_bytes / (1024 * 1024),
|
|
413
|
+
"full_memory_mb": total_full_bytes / (1024 * 1024),
|
|
414
|
+
"memory_saved_pct": (
|
|
415
|
+
(1 - total_state_bytes / max(total_full_bytes, 1)) * 100
|
|
416
|
+
),
|
|
417
|
+
"rank": self.rank,
|
|
418
|
+
"block_size": self.block_size,
|
|
419
|
+
"step_count": self._step_count,
|
|
420
|
+
}
|