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:
@@ -12,6 +12,9 @@
|
|||||||
# Secretos
|
# Secretos
|
||||||
.env
|
.env
|
||||||
|
|
||||||
|
# Config del IDE, específica de cada máquina
|
||||||
|
.vscode/
|
||||||
|
|
||||||
# Entorno
|
# Entorno
|
||||||
.venv/
|
.venv/
|
||||||
__pycache__/
|
__pycache__/
|
||||||
|
|||||||
Vendored
-3
@@ -1,3 +0,0 @@
|
|||||||
{
|
|
||||||
"ROS2.distro": "jazzy"
|
|
||||||
}
|
|
||||||
@@ -8,6 +8,9 @@ include:
|
|||||||
|
|
||||||
# El modelo char es diminuto: en la 2060 entra un micro-lote grande sin
|
# 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.
|
# acumulación. Sobrescribe el perfil de la placa solo en lo que corresponde.
|
||||||
|
train:
|
||||||
|
run_name: smoke-2060
|
||||||
|
|
||||||
hardware:
|
hardware:
|
||||||
micro_batch_size: 64
|
micro_batch_size: 64
|
||||||
grad_accum_steps: 1
|
grad_accum_steps: 1
|
||||||
|
|||||||
@@ -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