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:
@@ -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())
|
||||
Reference in New Issue
Block a user