Files
enlace/enlace/train/train.py
T
msaldain 03dadc93da
YAML / yaml (push) Failing after 1m58s
Python / calidad (push) Successful in 6s
Python / tests (push) Failing after 1m22s
Aplicar ruff format a todo el código
Cambio mecánico, sin efecto en el comportamiento: la suite pasa igual antes y
después. Va en un commit propio para no tapar los cambios con sentido.

Se agregan además dos flujos de verificación que corren en cada carga al
repositorio:

- YAML: yamllint para sintaxis y estilo, más la carga de cada config contra su
  esquema de pydantic. Son cosas distintas — un YAML puede ser sintácticamente
  perfecto y estar roto igual, con 'run_nombre' en vez de 'run_name'. Ese paso
  no instala torch: se verificó que la capa de configuración no lo importa, así
  que corre en segundos en vez de descargar dos gigas y medio de CUDA.
- Python: ruff check, ruff format --check y la suite completa con torch de CPU.

Las rutas ignoradas de .yamllint.yml van ancladas con barra inicial. Sin
anclar, 'data/' y 'runs/' excluían configs/data/ y configs/runs/ — siete
archivos, justo los que más importa revisar — y el linter pasaba en verde sin
haber mirado nada. Es el mismo defecto que ya había aparecido en .gitignore.
2026-07-28 07:25:51 -03:00

271 lines
9.8 KiB
Python

"""Bucle de entrenamiento.
No hay un solo hiperparámetro en este archivo: todo viene de la config. Lo que
sí vive acá es la mecánica que tiene que ser correcta pase lo que pase —
acumulación de gradiente, precisión mixta, recorte de norma, checkpointing
atómico y reanudación exacta.
python -m enlace.train.train configs/runs/smoke-2060.yaml
python -m enlace.train.train configs/runs/pretrain-2060.yaml --resume
"""
from __future__ import annotations
import json
import math
import random
import sys
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
from enlace.config.load import ConfigError, load_config
from enlace.config.schema import Config
from enlace.data.loaders import build_stream
from enlace.model.transformer import Transformer
from enlace.train import checkpoint
from enlace.train.backends import (
autocast_context,
build_grad_scaler,
describe_device,
setup_device,
)
from enlace.train.schedules import build_lr_fn
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def build_optimizer(model: torch.nn.Module, cfg: Config) -> torch.optim.AdamW:
"""AdamW con weight decay solo sobre las matrices.
Aplicar decay a normas y sesgos (tensores de una dimensión) degrada la
calidad sin ahorrar nada: son pocos parámetros y penalizarlos solo distorsiona
la normalización.
"""
decay, no_decay = [], []
for param in model.parameters():
if not param.requires_grad:
continue
(decay if param.dim() >= 2 else no_decay).append(param)
opt = cfg.train.optimizer
return torch.optim.AdamW(
[
{"params": decay, "weight_decay": opt.weight_decay},
{"params": no_decay, "weight_decay": 0.0},
],
lr=opt.lr,
betas=(opt.beta1, opt.beta2),
eps=opt.eps,
)
def flops_per_token(model: Transformer) -> float:
"""Estimación estándar de FLOPs por token en forward+backward.
6*N por las multiplicaciones de matrices y 12*L*H*Q*T por la atención, que
no depende del número de parámetros sino del largo de contexto.
"""
cfg = model.cfg
n = model.num_parameters(non_embedding=True)
attn = 12 * cfg.n_layer * cfg.n_head * cfg.head_dim * cfg.seq_len
return 6 * n + attn
@torch.no_grad()
def evaluate(model: torch.nn.Module, stream, cfg: Config, device: torch.device) -> float:
"""Loss promedio en el split de validación."""
model.eval()
total = 0.0
for _ in range(cfg.train.eval_batches):
x, y = stream.next_batch("val")
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
total += loss.item()
model.train()
return total / cfg.train.eval_batches
@torch.no_grad()
def sample(model: torch.nn.Module, stream, device: torch.device, n_tokens: int = 160) -> str:
"""Genera una muestra corta para inspección humana.
Leer estas muestras en cada checkpoint es la forma más rápida de detectar
que algo se rompió: las curvas de loss pueden verse bien mientras el modelo
produce basura.
"""
module = getattr(model, "_orig_mod", model)
start = torch.zeros((1, 1), dtype=torch.long, device=device)
out = module.generate(start, max_new_tokens=n_tokens, temperature=0.8, top_k=50)
model.train()
# Se descarta el token semilla: no es salida del modelo y ensucia la lectura.
generado = out[0, 1:].tolist()
decode = getattr(stream, "decode", None)
if decode is None:
return f"<sin decodificador> {generado[:16]}"
return decode(generado)
def train(cfg: Config, resume: bool = False) -> Path:
run_dir = Path(cfg.train.out_dir) / cfg.train.run_name
run_dir.mkdir(parents=True, exist_ok=True)
log_path = run_dir / "metrics.jsonl"
seed_everything(cfg.train.seed)
device = setup_device(cfg.hardware)
print(f"[enlace] dispositivo: {describe_device(cfg.hardware)}")
print(f"[enlace] perfil: {cfg.hardware.name} | dtype {cfg.hardware.dtype}")
stream = build_stream(cfg, device)
if stream.vocab_size != cfg.model.vocab_size:
raise ConfigError(
f"el corpus tiene vocab_size={stream.vocab_size} pero el modelo "
f"declara {cfg.model.vocab_size}. Son incompatibles."
)
model = Transformer(cfg.model, cfg.hardware.attention_backend).to(device)
n_params = model.num_parameters()
print(f"[enlace] modelo {cfg.model.name}: {n_params:,} parámetros")
optimizer = build_optimizer(model, cfg)
scaler = build_grad_scaler(cfg.hardware)
lr_at = build_lr_fn(cfg.train.schedule, cfg.train.optimizer.lr, cfg.train.max_steps)
start_step = 0
if resume:
ckpt_path = checkpoint.latest(run_dir)
if ckpt_path is None:
print(f"[enlace] --resume sin checkpoints en {run_dir}: se empieza de cero")
else:
meta = checkpoint.load(ckpt_path, model=model, optimizer=optimizer, scaler=scaler)
stream.load_state_dict(meta["stream"])
start_step = meta["step"]
print(f"[enlace] reanudado desde {ckpt_path.name} en el paso {start_step}")
if cfg.hardware.compile:
print("[enlace] compilando el modelo (la primera iteración tarda)...")
model = torch.compile(model) # type: ignore[assignment]
tokens_per_step = cfg.hardware.effective_batch_size * cfg.model.seq_len
fpt = flops_per_token(getattr(model, "_orig_mod", model))
model.train()
print(
f"[enlace] {cfg.train.max_steps} pasos x {tokens_per_step:,} tokens "
f"= {cfg.train.max_steps * tokens_per_step / 1e9:.2f}B tokens"
)
t_last = time.perf_counter()
for step in range(start_step, cfg.train.max_steps):
lr = lr_at(step)
for group in optimizer.param_groups:
group["lr"] = lr
optimizer.zero_grad(set_to_none=True)
loss_sum = 0.0
for _ in range(cfg.hardware.grad_accum_steps):
x, y = stream.next_batch("train")
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
# Se divide por los micro-pasos para que el gradiente acumulado
# sea el promedio y no la suma: si no, el LR efectivo dependería
# de grad_accum_steps y las curvas no serían comparables entre
# perfiles de hardware.
loss = loss / cfg.hardware.grad_accum_steps
scaler.scale(loss).backward()
loss_sum += loss.item()
grad_norm = float("nan")
if cfg.train.optimizer.grad_clip > 0:
# Hay que deshacer el escalado antes de medir la norma, o el recorte
# se aplicaría sobre gradientes inflados por el GradScaler.
scaler.unscale_(optimizer)
grad_norm = float(
torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.train.optimizer.grad_clip)
)
scaler.step(optimizer)
scaler.update()
if (step + 1) % cfg.train.log_every == 0:
if device.type == "cuda":
torch.cuda.synchronize()
now = time.perf_counter()
dt = (now - t_last) / cfg.train.log_every
t_last = now
record: dict[str, Any] = {
"step": step + 1,
"loss": round(loss_sum, 4),
"lr": lr,
"grad_norm": round(grad_norm, 4),
"tokens_per_s": round(tokens_per_step / dt),
"seconds_per_step": round(dt, 4),
}
if cfg.hardware.peak_tflops:
mfu = (fpt * tokens_per_step / dt) / (cfg.hardware.peak_tflops * 1e12)
record["mfu"] = round(mfu, 4)
with log_path.open("a") as fh:
fh.write(json.dumps(record) + "\n")
mfu_txt = f" mfu {record['mfu']:.1%}" if "mfu" in record else ""
print(
f"paso {step + 1:>7} | loss {loss_sum:.4f} | lr {lr:.2e} | "
f"|g| {grad_norm:.2f} | {record['tokens_per_s']:,} tok/s{mfu_txt}"
)
if (step + 1) % cfg.train.eval_every == 0:
val = evaluate(model, stream, cfg, device)
print(f"paso {step + 1:>7} | val_loss {val:.4f} | ppl {math.exp(min(val, 20)):.2f}")
with log_path.open("a") as fh:
fh.write(json.dumps({"step": step + 1, "val_loss": round(val, 4)}) + "\n")
t_last = time.perf_counter()
if cfg.train.sample_every and (step + 1) % cfg.train.sample_every == 0:
text = sample(model, stream, device)
(run_dir / "samples.txt").open("a").write(f"--- paso {step + 1} ---\n{text}\n\n")
print(f"muestra: {text[:120]!r}")
t_last = time.perf_counter()
if (step + 1) % cfg.train.checkpoint_every == 0 or (step + 1) == cfg.train.max_steps:
path = checkpoint.save(
run_dir / f"ckpt-{step + 1:08d}.pt",
config=cfg,
step=step + 1,
model=model,
optimizer=optimizer,
scaler=scaler,
stream_state=stream.state_dict(),
metrics={"loss": loss_sum},
)
print(f"[enlace] checkpoint -> {path.name}")
t_last = time.perf_counter()
return run_dir
def main() -> int:
args = [a for a in sys.argv[1:] if a != "--resume"]
resume = "--resume" in sys.argv[1:]
if not args:
print("uso: python -m enlace.train.train <config.yaml> [--resume] [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
train(cfg, resume=resume)
return 0
if __name__ == "__main__":
raise SystemExit(main())