03dadc93da
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.
271 lines
9.8 KiB
Python
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())
|