Verificación previa: validar una config contra el hardware real

Comprometer uno o dos días de cómputo para descubrir en la hora tres que el
lote no entra en memoria, o que fp16 diverge, es el desperdicio más caro del
proyecto. Esto responde en minutos: construye el modelo y el optimizador como
lo haría el entrenamiento y los ejercita con lotes sintéticos del tamaño
configurado, así que no necesita que el corpus exista.

Mide el pico de VRAM reservada (no la asignada: es la reservada la que hace
fallar la asignación), el rendimiento real convertido a horas de reloj, y la
tasa de pasos que el GradScaler descarta por inf/NaN — la señal directa de que
fp16 está perdiendo actualizaciones.

Aparte: .vscode/ sale del repo (config de IDE por máquina) y smoke-2060 fija
run_name, que sin eso los comandos de logs y de descarga pedían nombres
distintos para la misma corrida.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-28 00:48:59 -03:00
parent 36040d7cc8
commit 12bdc983e6
4 changed files with 229 additions and 3 deletions
+3
View File
@@ -12,6 +12,9 @@
# Secretos
.env
# Config del IDE, específica de cada máquina
.vscode/
# Entorno
.venv/
__pycache__/
-3
View File
@@ -1,3 +0,0 @@
{
"ROS2.distro": "jazzy"
}
+3
View File
@@ -8,6 +8,9 @@ include:
# El modelo char es diminuto: en la 2060 entra un micro-lote grande sin
# acumulación. Sobrescribe el perfil de la placa solo en lo que corresponde.
train:
run_name: smoke-2060
hardware:
micro_batch_size: 64
grad_accum_steps: 1
+223
View File
@@ -0,0 +1,223 @@
"""Validación de una config de entrenamiento contra el hardware real.
Comprometer uno o dos días de cómputo para descubrir en la hora tres que el
lote no entra en memoria, o que fp16 diverge, es el desperdicio más caro de
este proyecto. Esto responde antes, en un par de minutos y **sin necesitar el
corpus**: construye el modelo y el optimizador exactamente como los construiría
el entrenamiento, y los ejercita con lotes sintéticos del tamaño configurado.
Lo que se mide, y por qué cada cosa:
- **Pico de VRAM.** Es el que decide si la corrida entra. Se reporta el
reservado, no el asignado: el asignado ignora la fragmentación del caché de
PyTorch, y es lo reservado lo que hace fallar la asignación.
- **Rendimiento y MFU.** Convierte `max_steps` en horas de reloj. Sin esto,
"1 a 2 días" es una conjetura.
- **Estabilidad de fp16.** El GradScaler saltea los pasos donde detecta inf o
NaN. Unos pocos al principio son normales — está calibrando la escala —
pero una tasa sostenida significa que la corrida está perdiendo
actualizaciones y hay que bajar el LR o revisar la arquitectura.
python -m enlace.train.verificacion configs/runs/pretrain-2060.yaml
"""
from __future__ import annotations
import sys
import time
from dataclasses import dataclass
import torch
from enlace.config.load import ConfigError, load_config
from enlace.config.schema import Config
from enlace.model.transformer import Transformer
from enlace.train.backends import autocast_context, build_grad_scaler, describe_device, setup_device
from enlace.train.train import build_optimizer, flops_per_token
@dataclass
class Resultado:
ok: bool
titulo: str
detalle: str
def __str__(self) -> str:
marca = "OK " if self.ok else "FALLA"
return f" [{marca}] {self.titulo}: {self.detalle}"
def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
device = setup_device(cfg.hardware)
resultados: list[Resultado] = []
print(f"[enlace] dispositivo: {describe_device(cfg.hardware)}")
print(f"[enlace] perfil {cfg.hardware.name} | modelo {cfg.model.name} | dtype {cfg.hardware.dtype}")
model = Transformer(cfg.model, cfg.hardware.attention_backend).to(device)
optimizer = build_optimizer(model, cfg)
scaler = build_grad_scaler(cfg.hardware)
n_params = model.num_parameters()
tokens_por_paso = cfg.hardware.effective_batch_size * cfg.model.seq_len
fpt = flops_per_token(model)
if cfg.hardware.compile:
print("[enlace] compilando (la primera iteración tarda)...")
model = torch.compile(model) # type: ignore[assignment]
model.train()
if device.type == "cuda":
torch.cuda.reset_peak_memory_stats()
# Lotes sintéticos: para medir memoria y velocidad da igual el contenido,
# y así la verificación no depende de que el corpus ya exista.
generador = torch.Generator(device="cpu").manual_seed(cfg.train.seed)
def micro_lote():
datos = torch.randint(
0, cfg.model.vocab_size,
(cfg.hardware.micro_batch_size, cfg.model.seq_len + 1),
generator=generador,
)
return datos[:, :-1].to(device), datos[:, 1:].to(device)
tiempos: list[float] = []
perdidas: list[float] = []
salteados = 0
for paso in range(pasos):
if device.type == "cuda":
torch.cuda.synchronize()
t0 = time.perf_counter()
optimizer.zero_grad(set_to_none=True)
acumulada = 0.0
for _ in range(cfg.hardware.grad_accum_steps):
x, y = micro_lote()
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
loss = loss / cfg.hardware.grad_accum_steps
scaler.scale(loss).backward()
acumulada += loss.item()
if cfg.train.optimizer.grad_clip > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.train.optimizer.grad_clip)
# Si la escala baja tras step(), el GradScaler encontró inf/NaN y
# descartó la actualización. Es la señal directa de inestabilidad fp16.
escala_antes = scaler.get_scale()
scaler.step(optimizer)
scaler.update()
if scaler.get_scale() < escala_antes:
salteados += 1
if device.type == "cuda":
torch.cuda.synchronize()
dt = time.perf_counter() - t0
if paso >= warmup: # los primeros pasos incluyen compilación y autotune
tiempos.append(dt)
perdidas.append(acumulada)
# --- memoria ---
if device.type == "cuda":
pico_reservado = torch.cuda.max_memory_reserved() / 1024**3
total = torch.cuda.get_device_properties(0).total_memory / 1024**3
uso = pico_reservado / total
resultados.append(
Resultado(
uso < 0.92,
"VRAM",
f"pico reservado {pico_reservado:.2f} GB de {total:.1f} GB ({uso:.0%}). "
+ ("Entra con margen." if uso < 0.85 else
"Ajustado: un pico de fragmentación puede tirar la corrida." if uso < 0.92 else
"NO ENTRA con seguridad — bajá micro_batch_size o seq_len."),
)
)
# --- rendimiento ---
medio = sum(tiempos) / len(tiempos)
tok_s = tokens_por_paso / medio
mfu = (fpt * tokens_por_paso / medio) / (cfg.hardware.peak_tflops * 1e12) if cfg.hardware.peak_tflops else None
horas = cfg.train.max_steps * medio / 3600
tokens_totales = cfg.train.max_steps * tokens_por_paso
resultados.append(
Resultado(
True,
"Rendimiento",
f"{tok_s:,.0f} tok/s | {medio*1000:.0f} ms/paso"
+ (f" | MFU {mfu:.1%}" if mfu else ""),
)
)
resultados.append(
Resultado(
True,
"Proyección",
f"{cfg.train.max_steps:,} pasos x {tokens_por_paso:,} tok = "
f"{tokens_totales/1e9:.2f}B tokens en ~{horas:.1f} h ({horas/24:.1f} días)",
)
)
if mfu is not None:
resultados.append(
Resultado(
mfu > 0.15,
"MFU",
f"{mfu:.1%}{'razonable para esta placa' if mfu > 0.15 else 'bajo: revisá backend de atención, compile o tamaño de lote'}",
)
)
# --- estabilidad ---
finitas = all(p == p and abs(p) != float("inf") for p in perdidas)
resultados.append(
Resultado(finitas, "Loss finita", f"{len(perdidas)} pasos medidos, todas finitas" if finitas else "apareció NaN o inf")
)
if cfg.hardware.use_grad_scaler:
tasa = salteados / pasos
resultados.append(
Resultado(
tasa <= 0.5,
"Estabilidad fp16",
f"{salteados}/{pasos} pasos salteados por el GradScaler"
+ (" (normal al calibrar la escala)" if tasa <= 0.5
else " — tasa alta: la corrida perdería actualizaciones"),
)
)
# --- parámetros ---
resultados.append(
Resultado(True, "Modelo", f"{n_params:,} parámetros ({n_params/1e6:.1f}M)")
)
return resultados
def main() -> int:
args = sys.argv[1:]
if not args:
print("uso: python -m enlace.train.verificacion <config.yaml> [clave=valor ...]")
return 2
try:
cfg = load_config(args[0], args[1:])
except ConfigError as exc:
print(f"[enlace] {exc}", file=sys.stderr)
return 2
resultados = verificar(cfg)
print(f"\n[enlace] verificación previa de {args[0]}")
for r in resultados:
print(r)
fallas = [r for r in resultados if not r.ok]
print()
if fallas:
print(f"[enlace] {len(fallas)} verificación(es) fallaron — NO arranques la corrida larga.")
return 1
print("[enlace] configuración validada contra el hardware real.")
return 0
if __name__ == "__main__":
raise SystemExit(main())