90ac2a6582
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>
277 lines
9.9 KiB
Python
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())
|