Files
enlace/enlace/train/train.py
T
msaldain 90ac2a6582 Etapa 0: entorno, configuración validada, modelo y entrenador
Base del proyecto ENLACE: un modelo de lenguaje propio entrenado desde cero,
en español, para asistencia general y familiar. El plan completo está en
docs/PLAN.md.

Esta etapa establece el andamiaje y lo verifica de punta a punta:

- Configuración por capas (hardware × model × train × data) validada con
  pydantic. Ningún hiperparámetro vive en el código y una config inválida
  falla al arrancar, no a las tres horas de entrenamiento.
- Perfiles de hardware que aíslan el salto de GPU: la RTX 2060 (Turing) no
  soporta bfloat16 ni FlashAttention-2, así que entrena en float16 con
  GradScaler y backend mem_efficient; el perfil de la 5090 ya está escrito.
  backends.py valida el perfil contra la GPU real antes de empezar.
- Transformer decoder-only estilo Llama: RMSNorm, SwiGLU, RoPE, GQA,
  embeddings atados, QK-norm y z-loss. Los dos últimos son lo que mantiene
  estable el entrenamiento en float16.
- Entrenador con schedule WSD, acumulación de gradiente, precisión mixta,
  checkpointing atómico y reanudación exacta.
- Cargadores de datos con estado serializable: bytes para el smoke test y
  shards uint16 para el corpus real.

48 tests, entre ellos el crítico: reanudar desde un checkpoint reproduce los
pesos de una corrida ininterrumpida, parámetro por parámetro.

Verificado en CPU: 300 pasos sobre texto en español, loss 3.07 -> 1.63.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-27 23:04:45 -03:00

277 lines
9.9 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, map_location=device
)
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())