// M2 · lección 03
La receta de entrenamiento
AdamW, warmup y coseno, clipping y bf16, con el bucle entero de `rukh train` delante: el planificador de seis líneas, los grupos de decaimiento, la acumulación de gradiente ponderada por tokens, los checkpoints con procedencia y las dos configuraciones del repositorio.
Qué vas a construir
uv run rukh train --config configs/train/small.yaml, y los 608 líneas de src/rukh/train/ que hay
detrás. Al terminar tendrás tiny entrenado en 189 segundos y small en 42 minutos, con
sus veinte checkpoints, sus curvas en MLflow y la propiedad de la que depende todo lo demás: si el
proceso se corta, --resume continúa donde estaba en vez de empezar otra vez.
Un modelo bien construido no entrena solo. La receta que sigue es casi la misma que usaría cualquiera para un GPT pequeño, y el objeto de esta lección es que cada número tenga un motivo que sepas defender.
// Antes de empezarQué cuesta cada lab, y cuál puedes saltarte
| Lab | Ruta | Reloj | Deja | Atajo |
|---|---|---|---|---|
| tinyEntrenar tiny de principio a fin, para ver el bucle entero funcionar | cuesta máquina | 3 min de entrenamiento + ~10 de suite rápida | checkpoints/tiny/best.pt, artifacts/eval/tiny/ | rukh pull tiny |
| smallsmall entero, el primer modelo que juega por encima de 1200 Elo | cuesta máquina | 42 min + ~30 de las 160 partidas | checkpoints/small/best.pt y su fila en results.json | rukh pull small |
- cuesta máquina
- Tiempo real de GPU, red o motor. El reloj es el de la RTX 5090 de referencia.
Teoría justa: los cinco números de la receta
AdamWAdamWOptimizador Adam con el decaimiento de pesos desacoplado del gradiente: mantiene medias móviles del gradiente (β₁) y de su cuadrado (β₂) para dar a cada parámetro su propio paso, y resta aparte una fracción del peso. En Rukh: β = 0,9/0,95, weight decay 0,1 aplicado solo a las matrices. con β = 0,9 / 0,95. Adam guarda dos medias móviles por parámetro: la del gradiente (β₁ = 0,9, una especie de inercia) y la de su cuadrado (β₂), y divide la primera por la raíz de la segunda, de manera que cada parámetro recibe un paso normalizado por lo ruidoso que es su gradiente. El valor por defecto de β₂ en PyTorch es 0,999, que promedia sobre unos mil pasos; en modelos de lenguaje se baja a 0,95 (unos veinte pasos) porque la escala de los gradientes cambia deprisa al principio y una media demasiado larga reacciona tarde. Es el valor de GPT-3 y el de nanoGPT.
Weight decay 0,1, solo en las matrices. El decaimiento de pesos resta en cada paso una fracción
del propio peso: empuja todo hacia cero salvo que el gradiente lo sostenga. En las matrices —la
atención, el MLP, los embeddings— eso es regularización clásica. En los parámetros de LayerNorm y
en los sesgos, no.
El decaimiento es un alquiler que se cobra por peso: cada parámetro paga cada paso una fracción de
lo que vale, y solo se queda grande el que el gradiente sostiene porque de verdad hace falta. Lo
que no se usa se va encogiendo hasta desaparecer. A las ganancias del LayerNorm no se les cobra
porque su tamaño no es capacidad: es un ajuste de escala, y cobrarles alquiler es apagar el canal.
WarmupWarmupArranque en el que la tasa de aprendizaje sube linealmente desde cero durante los primeros pasos (1 000 en `small`) antes de empezar a decaer. Evita que Adam dé pasos enormes mientras sus medias móviles todavía se estiman con un puñado de gradientes —en el paso 1, con uno solo—, que es cuando un modelo recién inicializado se rompe: no con un `NaN`, sino cayendo en predecir siempre las jugadas más frecuentes del corpus. de 1 000 pasos y después coseno. La tasa de aprendizaje empieza en cero, sube en línea recta hasta 6e-4 en el paso 1 000 y desde ahí baja siguiendo medio coseno hasta el 10 % de ese valor en el paso 20 000, donde se queda.
Qué pasa sin warmup. En el paso 1, Adam tiene una estimación del segundo momento hecha con un único
gradiente; el cociente que calcula es casi arbitrario —es calcular tu velocidad media en un viaje
cuando llevas recorrido el primer metro: el número existe, pero no significa nada todavía—, y como
el paso de Adam está normalizado, un parámetro puede moverse tanto como la tasa de aprendizaje de
golpe. El warmup es no fiarse de esa media hasta que hay kilómetros detrás. Sobre una red recién
inicializada, con 24 escrituras en la corriente residual, eso basta para que las activaciones
exploten, el softmax de la atención sature y el modelo caiga en un mínimo tonto: predecir siempre
las jugadas más frecuentes del corpus. La curva no revienta con un NaN espectacular; se queda
plana en una pérdida mediocre y no baja nunca. El warmup cuesta el 5 % de los pasos y elimina ese
riesgo entero.
La bajada en coseno tiene una lógica parecida por el otro lado: al final del entrenamiento el modelo está cerca de un mínimo y pasos grandes solo lo sacan de él. El suelo del 10 % evita que los últimos miles de pasos no hagan nada.
Clippingrecorte de gradienteReescalar el vector de gradientes cuando su norma global supera un umbral (1,0 en Rukh), antes de cada paso del optimizador. Es un seguro contra un lote patológico: sin él, una secuencia rara puede mover los pesos lo suficiente para tirar horas de entrenamiento. El `grad_norm` registrado en MLflow pegado al umbral paso tras paso significa que la tasa de aprendizaje es alta. a 1,0. Antes de cada optimizer.step() se calcula la norma global del gradiente y, si
pasa de 1,0, se reescala el vector entero para que valga 1,0. Es un seguro contra un lote raro: sin
él, una única secuencia patológica puede mover los pesos lo suficiente como para tirar horas de
entrenamiento. El grad_norm se registra en MLflow, y mirarlo es la mejor forma de saber si el
entrenamiento está sano: debe bajar y estabilizarse, no dar picos. Es el limitador de velocidad de
un coche de alquiler: no cambia cómo conduces mientras vas por debajo, y el día que pisas de más
porque la carretera te engañó, es lo que evita el accidente. Si salta en cada curva —el grad_norm
pegado a 1,0 paso tras paso—, el problema no es el limitador, es que vas demasiado deprisa para esa
carretera: la tasa de aprendizaje es alta.
bf16bfloat16Formato de coma flotante de 16 bits con los mismos 8 bits de exponente que fp32 y solo 7 de mantisa: pierde precisión pero conserva el rango, así que un gradiente pequeño no se va a cero. Por eso entrena sin escalado de pérdida, al contrario que fp16, cuyo exponente de 5 bits desborda por abajo., y por qué no fp16. Los dos son formatos de 16 bits y los dos
doblan el rendimiento en una GPU moderna. La diferencia está en cómo reparten los bits. fp16 usa 5
para el exponente y 10 para la mantisa: mucha precisión, poco rango, con el número normal más
pequeño alrededor de 6·10⁻⁵. Los gradientes de una red profunda viven justo ahí abajo, así que en
fp16 se van a cero: por eso entrenar en fp16 obliga a un GradScaler que multiplica la pérdida por
un factor grande, lo baja si aparecen infinitos, y de vez en cuando te regala un NaN a las tres
de la mañana. bf16 usa 8 bits de exponente —el mismo rango que fp32— y 7 de mantisa: menos
precisión por número, pero nada se desborda por abajo. No hace falta escalar nada. Son dos reglas
del mismo precio: fp16 es una regla fina y corta —marca las décimas de milímetro, pero solo llega
a 20 centímetros—, bf16 es gruesa y larga —milímetros enteros, pero llega hasta donde llega una
de metro—. Para un gradiente, que puede ser minúsculo o enorme y casi nunca necesita la tercera
cifra, la larga es la que sirve. En la RTX 5090 (sm_120) bf16 es nativo, así que la elección es
gratis. Los pesos maestros siguen en fp32 y el optimizador acumula en fp32; bf16 es solo la
precisión de las operaciones dentro del autocast.
schedule.py: veintiocho líneas, y las tres primeras deciden la corrida
Empezamos por el fichero más pequeño del módulo, que es además el que separa una corrida que converge de una que se queda plana veinte mil pasos.
"""Learning-rate schedule: a linear warmup followed by a cosine decay to a floor."""
from __future__ import annotations
import mathfrom typing import TYPE_CHECKING
if TYPE_CHECKING: # pragma: no cover - only for type checkers from rukh.train.loop import TrainConfigEl TYPE_CHECKING no es un adorno: schedule.py necesita el tipo de TrainConfig para anotarse, y
loop.py importa lr_at. Importarlo de verdad sería un ciclo. Bajo TYPE_CHECKING el import no
ocurre en tiempo de ejecución y el verificador de tipos sí lo ve, que es exactamente lo que hace
falta.
def lr_at(step: int, cfg: TrainConfig) -> float: """Learning rate for ``step`` (0-based).
It rises linearly from 0 to ``cfg.lr`` over the first ``cfg.warmup`` steps, then follows a cosine down to ``cfg.min_lr_ratio * cfg.lr`` at ``cfg.max_steps`` and stays there for any later step (so a resumed run past ``max_steps`` never gets a negative or rising rate). """ if step < 0: raise ValueError(f"step must be non-negative, got {step}") floor = cfg.lr * cfg.min_lr_ratio if step < cfg.warmup: return cfg.lr * step / cfg.warmup span = cfg.max_steps - cfg.warmup if span <= 0: return floor progress = min(1.0, (step - cfg.warmup) / span) return floor + 0.5 * (1.0 + math.cos(math.pi * progress)) * (cfg.lr - floor)La rampa es una regla de tres y el coseno son dos líneas. Lo que merece un párrafo son las dos guardas:
step < 0es un error, no un cero. Una tasa negativa no existe, y si un llamante calcula mal un paso lo que quieres es enterarte ahí y no cuarenta minutos después.span <= 0devuelve el suelo. Ocurre cuandowarmup >= max_steps, que suena absurdo hasta que alguien lanza--max-steps 300sobre una configuración conwarmup: 500para probar algo rápido. Sin esa rama, la división por cero. Y pasadomax_stepselmin(1.0, …)deja la tasa clavada en el suelo, de modo que una corrida reanudada que se pase nunca recibe una tasa negativa ni creciente.
"""Tests for rukh.train.schedule: linear warmup then cosine decay to a floor."""
from __future__ import annotations
import pytest
from rukh.train import TrainConfig, lr_at
pytestmark = pytest.mark.unit
CFG = TrainConfig(lr=6e-4, min_lr_ratio=0.1, warmup=100, max_steps=1000)
def test_warmup_is_linear_and_reaches_the_peak() -> None: assert lr_at(0, CFG) == 0.0 assert lr_at(50, CFG) == pytest.approx(CFG.lr / 2) assert lr_at(100, CFG) == pytest.approx(CFG.lr) rising = [lr_at(step, CFG) for step in range(101)] assert rising == sorted(rising) deltas = {round(b - a, 12) for a, b in zip(rising, rising[1:], strict=False)} assert len(deltas) == 1El test del warmup no comprueba tres puntos: comprueba que los 101 primeros valores están ordenados y que todos los incrementos son el mismo número (el conjunto de deltas redondeadas tiene un solo elemento). Eso es lo que significa «lineal», y es lo que distingue una rampa de cualquier otra curva creciente que pase por los tres puntos que se te ocurra mirar.
def test_cosine_decays_to_the_floor_and_stays_there() -> None: floor = CFG.lr * CFG.min_lr_ratio assert lr_at(1000, CFG) == pytest.approx(floor) assert lr_at(5000, CFG) == pytest.approx(floor) falling = [lr_at(step, CFG) for step in range(100, 1001, 10)] assert falling == sorted(falling, reverse=True) assert all(value >= floor - 1e-12 for value in falling) halfway = lr_at(550, CFG) assert halfway == pytest.approx(floor + 0.5 * (CFG.lr - floor))
def test_without_warmup_the_first_step_is_the_peak() -> None: cfg = CFG.model_copy(update={"warmup": 0}) assert lr_at(0, cfg) == pytest.approx(cfg.lr)
def test_degenerate_schedule_returns_the_floor() -> None: cfg = CFG.model_copy(update={"warmup": 1000, "max_steps": 1000}) assert lr_at(1000, cfg) == pytest.approx(cfg.lr * cfg.min_lr_ratio) with pytest.raises(ValueError, match="non-negative"): lr_at(-1, CFG)
def test_config_rejects_unknown_keys() -> None: with pytest.raises(ValueError): TrainConfig(learning_rate=1e-3) # type: ignore[call-arg]Y los otros cuatro cubren el otro extremo: que el coseno baje de forma monótona, que el punto medio
del descenso valga exactamente la media entre el pico y el suelo (la propiedad del coseno que
verifica que la fórmula es la que crees), que pasado max_steps no siga bajando, y los dos casos
degenerados. El último no es del planificador sino de la configuración: TrainConfig(learning_rate=…)
tiene que explotar, porque learning_rate no es el nombre del campo y un YAML con esa clave
entrenaría con 6e-4 sin decir nada.
loop.py: el bucle
Trescientas sesenta y tres líneas, y la mitad son la contabilidad que hace que una corrida sea reproducible seis meses después. Empezamos por la cabecera, que es un resumen de las decisiones.
"""The training loop: AdamW with a cosine schedule, bf16, checkpoints and MLflow logging.
The recipe follows ``docs/spec/02`` (component 1): AdamW (0.9/0.95, weight decay 0.1 appliedonly to matrices), learning rate 6e-4 with 1 000 warmup steps and a cosine decay, an effectivebatch of ``batch_size * grad_accum`` sequences, bf16 autocast on CUDA, gradient clipping at 1.0and optional ``torch.compile``. Everything the run needs to be reproducible (config, seed, gitSHA, vocabulary hash, data manifest hash) goes to MLflow and into every checkpoint.
Resuming is meant to be indistinguishable from never having stopped: the batch stream is woundforward past the windows the first half of the run already saw, and the MLflow run id travels inthe checkpoint so the curve carries on in the same run instead of starting a second one.
Two throughput numbers are logged because they answer different questions: ``tokens_per_s``counts every position in the window (what the GPU actually processed, comparable across runs)and ``real_tokens_per_s`` counts only the non-``<pad>`` targets (what the model learned from).The training loss is accumulated token-weighted rather than as a mean of means, so amicro-batch with fewer real tokens does not count as much as a full one."""Léela entera: es la especificación de la receta, escrita donde no se pierde. Los dos párrafos que importan son el de reanudar —«que sea indistinguible de no haber parado nunca» es un requisito, no una intención— y el de las dos cifras de rendimiento.
from __future__ import annotations
import loggingimport mathimport osimport timefrom collections.abc import Iteratorfrom contextlib import nullcontextfrom datetime import UTC, datetimefrom pathlib import Pathfrom typing import Any, Literal
import torchfrom torch import nnfrom torch.utils.data import DataLoader
from rukh import pathsfrom rukh.config import BaseConfigfrom rukh.models import DecoderConfig, MoveDecoder, presetfrom rukh.tokenize.loader import IGNORE_INDEX, PackedDataset, make_loaderfrom rukh.train.checkpoint import ( BEST_NAME, load_checkpoint, read_manifest_sha, read_vocab_hash, restore, save_checkpoint, step_name,)from rukh.train.schedule import lr_at
log = logging.getLogger(__name__)
Batch = tuple[torch.Tensor, torch.Tensor]class TrainConfig(BaseConfig): """Everything one training run needs; unknown keys in the YAML are an error."""
preset: Literal["tiny", "small", "medium"] = "small" model: DecoderConfig | None = None # overrides the preset when given tokens_dir: str = "data/tokens/uci" block: int = 200 batch_size: int = 64 grad_accum: int = 4 # 256 effective sequences lr: float = 6e-4 min_lr_ratio: float = 0.1 warmup: int = 1000 max_steps: int = 20000 weight_decay: float = 0.1 betas: tuple[float, float] = (0.9, 0.95) grad_clip: float = 1.0 precision: Literal["bf16", "fp32"] = "bf16" compile: bool = True eval_every: int = 500 eval_batches: int = 50 ckpt_every: int = 1000 out_dir: str = "checkpoints" seed: int = 42 run_name: str | None = None unique_run_name: bool = True """Append a timestamp to ``run_name``: a second run must not overwrite the ``step-*.pt`` series the ``TrainingReplay`` of the course reads.""" workers: int = 0 # DataLoader workers; 0 keeps everything in the main process log_every: int = 10 # optimizer steps between training metrics
def decoder(self) -> DecoderConfig: """The decoder config of this run: the preset (or ``model``) with ``block`` applied.""" base = self.model if self.model is not None else preset(self.preset) return base.model_copy(update={"block": self.block})TrainConfig es la receta entera, y como hereda de BaseConfig una clave desconocida en el YAML es
un error. Los valores por defecto son los de small. Tres campos que no son obvios:
grad_accum: 4, conbatch_size: 64, son 256 secuencias efectivas. La acumulación de gradienteacumulación de gradienteEn Rukh, partir el lote que no cabe en la GPU en varias pasadas que sí caben, acumulando sus gradientes antes de un solo paso del optimizador. `grad_accum: 4` con `batch_size: 64` son 256 secuencias efectivas; la pérdida de cada micro-lote se divide por `grad_accum` para que los gradientes sumados equivalgan a una media y el tamaño del paso no dependa de cuántos micro-lotes se acumulen. no es un lote más grande gratis: es partir el lote que no cabe en la GPU en cuatro pasadas que sí caben.unique_run_name: Trueañade una marca de tiempo al nombre. Sin ella, una segunda corrida de la misma configuración sobreescribiría losstep-*.ptde la primera, que son exactamente los que lee la islaTrainingReplayde este módulo. El docstring del campo lo dice ahí mismo, que es donde se lee.model: DecoderConfig | Nonepermite describir un modelo a mano en el YAML en vez de usar un preset. Es lo que hace queconfigs/train/*.yamlno tenga que crecer un preset por experimento.
Y decoder() es la única línea del fichero que toca la forma del modelo: coge el preset (o el
model explícito) y le aplica el block de la corrida. El contexto vive en dos sitios porque son
dos cosas distintas —la ventana del modelo y la longitud de las ventanas que empaqueta el
cargador—, y esta línea es la que garantiza que coincidan.
def pick_device() -> str: """``RUKH_DEVICE`` if set, else CUDA when available, else CPU.""" env = os.environ.get("RUKH_DEVICE") if env: return env return "cuda" if torch.cuda.is_available() else "cpu"
def param_groups(model: nn.Module, weight_decay: float) -> list[dict[str, Any]]: """Decoupled weight decay on matrices only: norms and biases are left alone.""" decay = [p for p in model.parameters() if p.requires_grad and p.dim() >= 2] no_decay = [p for p in model.parameters() if p.requires_grad and p.dim() < 2] return [ {"params": decay, "weight_decay": weight_decay}, {"params": no_decay, "weight_decay": 0.0}, ]pick_device mira primero RUKH_DEVICE, y esa variable de entorno es lo que permite a los tests
forzar CPU en una máquina con GPU sin tocar ninguna configuración.
param_groups son las seis líneas del decaimiento de pesos, y la regla es literalmente «tensores de
dos dimensiones o más». Un LayerNorm tiene una ganancia por canal inicializada a 1; empujarla hacia
cero es apagar el canal, no regularizarlo. Un sesgo tiene una dimensión y desplaza, no escala. Son
unos pocos miles de parámetros de los 39 millones, y meterlos en el grupo equivocado degrada el
entrenamiento en silencio, que es la peor forma de degradarlo.
def forever(loader: DataLoader[Batch]) -> Iterator[Batch]: """Repeat a loader for as many steps as the schedule asks for.""" while True: yield from loaderCuatro líneas que convierten un DataLoader finito en un flujo infinito. El bucle cuenta pasos, no
épocas, así que necesita que los lotes no se acaben; yield from reinicia el cargador —y con él su
barajado— cada vez que se agota. En la tirada de small eso ocurre 4,3 veces, y esa cifra es la que
la última lección del módulo convierte en un diagnóstico.
def maybe_compile( model: MoveDecoder, enabled: bool, sample: torch.Tensor | None = None) -> nn.Module: """``torch.compile`` the model when asked; a failure is a warning, never a stopped run.
``torch.compile`` is lazy: on Windows without MSVC it only fails when the first forward reaches Inductor. ``sample`` (one batch of the training shape) forces that compilation here, where it can still fall back to eager, and Dynamo's own error suppression covers any later recompilation for a different shape. """ if not enabled: return model try: import torch._dynamo as dynamo
dynamo.config.suppress_errors = True compiled = torch.compile(model) if sample is not None: _, loss = compiled(sample, sample) if loss is not None: loss.backward() model.zero_grad(set_to_none=True) return compiled except Exception as exc: # noqa: BLE001 - compilation backends fail in many ways log.warning("torch.compile is unavailable, training eagerly: %s", exc) model.zero_grad(set_to_none=True) return modeltorch.compile es perezoso: no compila al llamarlo, compila la primera vez que una pasada llega a
Inductor. En Windows sin MSVC (o sin Triton) eso significa que el fallo aparece a mitad de la
corrida, cuando ya has invertido media hora. Forzar la compilación aquí con un lote de la forma de
entrenamiento hace que el fallo ocurra donde todavía se puede caer a modo eager.
Y la caída es un aviso, no una excepción: log.warning y se entrena igual. La máquina de referencia
entrena small entero en eager por exactamente esto, a 422 000 tokens por segundo, y lo dice en vez
de fingir que la bandera funcionó.
@torch.no_grad()def evaluate( model: nn.Module, loader: DataLoader[Batch], batches: int, device: torch.device, autocast: Any = None,) -> tuple[float, float]: """Validation loss and top-1 next-token accuracy over at most ``batches`` batches.
The loss is token-weighted: each batch's mean is weighted by the number of non-``<pad>`` targets it had, so a short last batch does not count as much as a full one. """ was_training = model.training model.eval() loss_sum = 0.0 weighted = 0 hits = 0 counted = 0 for index, (x, y) in enumerate(loader): if index >= batches: break x, y = x.to(device), y.to(device) with autocast if autocast is not None else nullcontext(): logits, loss = model(x, y) mask = y != IGNORE_INDEX tokens = int(mask.sum()) if loss is not None and torch.isfinite(loss) and tokens: loss_sum += loss.float().item() * tokens weighted += tokens hits += int((logits.argmax(dim=-1) == y)[mask].sum()) counted += tokens if was_training: model.train() return (loss_sum / weighted if weighted else math.nan, hits / counted if counted else 0.0)La evaluación, y tres decisiones dentro:
was_trainingen vez demodel.train()al final. La función se llama desde el bucle (donde el modelo está entrenando) y desde scripts (donde no). Restaurar el estado que había es lo que impide que evaluar un checkpoint lo deje silenciosamente en modo entrenamiento.- Un número fijo de lotes. 50 lotes de 64 × 200 son 640 000 tokens, y esta medida se satura mucho antes de que haga falta un mes entero de partidas de validación. Esa observación es la que la última lección del módulo convierte en 238 millones de tokens más de material de entrenamiento.
- La pérdida ponderada por tokens, no como media de medias: un lote con más relleno no debe
pesar lo mismo que uno lleno. Y el
torch.isfinite(loss)filtra el lote degenerado —todo<pad>— que la lección anterior dejó devolviendoNaN.
Fíjate también en (logits.argmax(dim=-1) == y)[mask]: el top-1 se cuenta solo sobre las posiciones
que no son relleno. Contarlo sobre todas daría una cifra altísima y sin sentido, porque acertar
<pad> es gratis.
def run_dir(cfg: TrainConfig, resume: Path | None = None) -> Path: """Where this run writes: the resumed run's folder, or ``out_dir/<run name>``.
The name carries a timestamp unless ``unique_run_name`` is off, because two runs of the same config would otherwise write the same ``step-*.pt`` files and the second would quietly overwrite the checkpoint series of the first. """ if resume is not None: return Path(resume).resolve().parent name = cfg.run_name or cfg.preset if cfg.unique_run_name: name = f"{name}-{datetime.now(UTC):%Y%m%d-%H%M%S}" return paths.resolve(cfg.out_dir) / name
def skip_batches(batches: Iterator[Batch], count: int) -> int: """Wind the batch stream forward ``count`` batches and return how many were skipped.
A resumed run must not start again at the first window of the first epoch: it would train twice on the same games while the schedule believes it is halfway. Windows are memmap slices, so winding forward is cheap compared with a step, and it is logged because it is not free. """ if count <= 0: return 0 started = time.perf_counter() for _ in range(count): next(batches) log.info("skipped %d batches in %.1f s to resume", count, time.perf_counter() - started) return countrun_dir decide dónde escribe la corrida y skip_batches es la mitad menos obvia de reanudar. Una
corrida reanudada que empezara otra vez por la primera ventana de la primera época entrenaría dos
veces sobre las mismas partidas mientras el planificador cree que va por la mitad. Adelantar el flujo
es barato —las ventanas son rodajas de un memmap— y se registra, porque no es gratis.
def train(cfg: TrainConfig, resume: Path | None = None, device: str | None = None) -> Path: """Train a ``MoveDecoder`` and return the path of the final checkpoint.""" if min(cfg.max_steps, cfg.grad_accum, cfg.batch_size) < 1: raise ValueError("max_steps, grad_accum and batch_size must be positive") torch.manual_seed(cfg.seed) where = torch.device(device or pick_device()) tokens_dir = paths.resolve(cfg.tokens_dir) model_cfg = cfg.decoder()
train_set = PackedDataset(tokens_dir / "train", block=cfg.block) val_set = PackedDataset(tokens_dir / "val", block=cfg.block) if train_set.info.vocab_size != model_cfg.vocab_size: raise ValueError( f"{tokens_dir / 'train'} has vocab_size {train_set.info.vocab_size}, " f"the model expects {model_cfg.vocab_size}" ) train_loader = make_loader(train_set, cfg.batch_size, seed=cfg.seed, workers=cfg.workers) val_loader = make_loader( val_set, cfg.batch_size, seed=cfg.seed, workers=0, shuffle=False, drop_last=False ) if not len(train_loader): raise ValueError(f"{tokens_dir / 'train'} has fewer than {cfg.batch_size} windows")La entrada de train. La comprobación que vale su peso es la del vocabulario: si el paquete de
tokens dice 2 030 y el modelo espera otra cosa, el entrenamiento funcionaría —los índices son
enteros y caben— y produciría un modelo que habla otro idioma. Ese error no se detecta después.
model = MoveDecoder(model_cfg).to(where) optimizer = torch.optim.AdamW(param_groups(model, cfg.weight_decay), lr=cfg.lr, betas=cfg.betas) start_step = 0 best_val = math.inf run_id: str | None = None if resume is not None: payload = load_checkpoint(resume, map_location=where) start_step = restore(payload, model, optimizer) recorded = payload.get("best_val") best_val = float(recorded) if isinstance(recorded, int | float) else best_val previous = payload.get("run_id") run_id = str(previous) if isinstance(previous, str) and previous else None log.info("resumed %s at step %d (mlflow run %s)", resume, start_step, run_id or "new") use_bf16 = cfg.precision == "bf16" and where.type == "cuda" autocast = ( torch.autocast(device_type="cuda", dtype=torch.bfloat16) if use_bf16 else nullcontext() ) warmup = torch.ones((cfg.batch_size, cfg.block), dtype=torch.long, device=where) with autocast: # compile under the same precision the loop will use runnable = maybe_compile(model, cfg.compile, warmup) runnable.train()
out_dir = run_dir(cfg, resume) out_dir.mkdir(parents=True, exist_ok=True) vocab_hash = read_vocab_hash(tokens_dir / "train") manifest_sha = read_manifest_sha(paths.data_dir() / "raw" / "manifest.json") tokens_per_step = cfg.batch_size * cfg.grad_accum * cfg.block batches = forever(train_loader) skip_batches(batches, start_step * cfg.grad_accum) final = out_dir / step_name(cfg.max_steps)El warmup de la línea 254 no es el warmup de la tasa de aprendizaje: es el lote de unos con el que
se fuerza la compilación, y va dentro del mismo autocast con el que va a correr el bucle. Si se
compilara en fp32 y se ejecutara en bf16, la primera pasada real recompilaría, que es justo lo que
este bloque existe para descubrir pronto.
from rukh.tracking import git_sha, start_run
params = { **cfg.model_dump(mode="json"), "vocab_hash": vocab_hash, "data_manifest_sha": manifest_sha, "device": str(where), "num_params": model.num_params(), } with start_run(out_dir.name, params, tags={"preset": cfg.preset}, run_id=run_id) as run: log.info("run %s in %s on %s", run.info.run_id, out_dir, where) this_run = str(run.info.run_id)
def save(path: Path, step: int) -> Path: return save_checkpoint( path, step=step, model=model, optimizer=optimizer, cfg=cfg.model_dump(mode="json"), model_cfg=model_cfg.model_dump(mode="json"), vocab_hash=vocab_hash, data_manifest_sha=manifest_sha, git_sha=git_sha(), best_val=None if math.isinf(best_val) else best_val, run_id=this_run, )Todo lo que hace falta para identificar una corrida seis meses después: la configuración entera, el
hash del vocabulario, el SHA del manifiesto de datos, el dispositivo y el número de parámetros. Y
save es una clausura porque necesita this_run, que solo existe dentro del with.
clock = time.perf_counter() for step in range(start_step, cfg.max_steps): lr = lr_at(step, cfg) for group in optimizer.param_groups: group["lr"] = lr optimizer.zero_grad(set_to_none=True) # Token-weighted, on the device: one synchronisation per step instead of one per # micro-batch, and a micro-batch with fewer real tokens weighs less in the mean. loss_sum = torch.zeros((), device=where) real_tokens = torch.zeros((), device=where) for _ in range(cfg.grad_accum): x, y = next(batches) x, y = x.to(where), y.to(where) with autocast: _, loss = runnable(x, y) assert loss is not None (loss / cfg.grad_accum).backward() tokens = (y != IGNORE_INDEX).sum() loss_sum += loss.detach().float() * tokens real_tokens += tokens grad_norm = float(nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)) optimizer.step()El corazón. Veintidós líneas, y tres detalles que se equivocan con facilidad:
(loss / cfg.grad_accum).backward(). Es lo que hace que los cuatro gradientes acumulados equivalgan a una media y no a una suma. Sin esa división, el tamaño del paso dependería de cuántos micro-lotes se acumulen, y cambiargrad_accumpara que quepa en otra GPU cambiaría la receta.- La pérdida se acumula en la GPU, en dos tensores escalares, y solo se lee con
.item()al registrar. Acumular en Python costaría una sincronización por micro-lote en vez de una por paso. clip_grad_norm_(model.parameters(), …), no la derunnable. Contorch.compile,runnablees un envoltorio; los parámetros son los mismos objetos, pero pedirlos al modelo original es lo que garantiza que el recorte se aplica a la norma global de todos ellos y no a la de un subconjunto.
done = step + 1 last = done == cfg.max_steps if done % cfg.log_every == 0 or last: counted = float(real_tokens.item()) total = float(loss_sum.item()) / counted if counted else math.nan elapsed = max(time.perf_counter() - clock, 1e-9) steps = min(cfg.log_every, done - start_step) log_metrics( { "train/loss": total, "lr": lr, "grad_norm": grad_norm, "tokens_per_s": tokens_per_step * steps / elapsed, "real_tokens_per_s": counted * steps / elapsed, }, step=done, ) clock = time.perf_counter() if done % cfg.eval_every == 0 or last: val_loss, top1 = evaluate( runnable, val_loader, cfg.eval_batches, where, autocast if use_bf16 else None ) log_metrics({"val/loss": val_loss, "val/top1": top1}, step=done) log.info("step %d val/loss %.4f val/top1 %.4f", done, val_loss, top1) if not math.isnan(val_loss) and val_loss < best_val: best_val = val_loss save(out_dir / BEST_NAME, done) clock = time.perf_counter() if done % cfg.ckpt_every == 0 or last: final = save(out_dir / step_name(done), done) clock = time.perf_counter() return finalLa contabilidad. done = step + 1 para que «cada 500 pasos» signifique lo que parece, y el or last
en las tres condiciones para que la última iteración siempre registre, evalúe y guarde aunque
max_steps no sea múltiplo de nada.
El best.pt se escribe solo cuando la pérdida de validación mejora, así que al acabar tienes dos
cosas distintas: el último checkpoint y el mejor. En la tirada de small son el mismo, y eso es una
medida, no una suposición: la validación seguía bajando en el paso 20 000.
Los tres clock = time.perf_counter() son el detalle aburrido que hace que tokens_per_s signifique
algo: sin ellos, el tiempo de una evaluación o de escribir un checkpoint de 150 MB entraría en la
medida de la velocidad de entrenamiento.
def log_metrics(metrics: dict[str, float], step: int) -> None: """Send metrics to the active MLflow run; a tracking failure never stops training.""" import mlflow
clean = {key: value for key, value in metrics.items() if not math.isnan(value)} if not clean: return try: mlflow.log_metrics(clean, step=step) except Exception as exc: # noqa: BLE001 - tracking is not worth a lost run log.warning("could not log metrics at step %d: %s", step, exc)Y el registro, con la regla que se aplica a todo el proyecto: perder la traza nunca cuesta una
corrida. Un MLflow caído, un disco lleno o una métrica NaN producen un aviso y el entrenamiento
sigue. El filtro de NaN no es cosmético: MLflow rechaza el valor y la excepción, sin el try,
tiraría la corrida en el primer lote raro.
checkpoint.py: qué es un checkpoint y qué no
"""Checkpoints: what a run writes so it can be resumed, evaluated and published.
A payload holds only plain Python values and tensors, so it can be read back with``torch.load(..., weights_only=True)``. Provenance (``vocab_hash``, ``data_manifest_sha``,``git_sha``) is recorded on purpose: weights whose data or vocabulary cannot be identified arenot reproducible."""
from __future__ import annotations
import hashlibimport jsonimport randomfrom collections.abc import Mappingfrom pathlib import Pathfrom typing import Any
import numpy as npimport torchfrom torch import nn
from rukh.models import DecoderConfig, MoveDecoder
BEST_NAME = "best.pt"TIED_HEAD = "lm_head.weight""""Tied to ``tokens.weight``; a published state dict leaves it out and it is re-tied on load."""«Solo valores de Python y tensores» tiene una consecuencia concreta: el fichero se puede leer con
torch.load(..., weights_only=True), que es el modo que no ejecuta código al deserializar. Un
checkpoint es un fichero que la gente se descarga; un pickle arbitrario es una ejecución remota
esperando a ocurrir.
def step_name(step: int) -> str: """File name of the checkpoint written after ``step`` optimizer steps.""" return f"step-{step}.pt"
def read_vocab_hash(tokens_dir: Path) -> str | None: """``vocab_hash`` from a pack ``meta.json``, or None when it cannot be read.""" try: meta = json.loads((Path(tokens_dir) / "meta.json").read_text(encoding="utf-8")) except (OSError, ValueError): return None value = meta.get("vocab_hash") return str(value) if isinstance(value, str) else None
def read_manifest_sha(manifest: Path) -> str | None: """SHA-256 of a data manifest, or None when the file is absent or unreadable.
Training must never depend on ``data/`` being present: a checkpoint trained from a pack copied elsewhere simply records ``None``. """ try: payload = Path(manifest).read_bytes() except OSError: return None return hashlib.sha256(payload).hexdigest()Las dos funciones de procedencia, y las dos devuelven None en vez de fallar. Es deliberado y está
escrito: entrenar desde un paquete de tokens copiado a otra máquina, donde no hay data/, tiene que
funcionar; lo que no puede pasar es que el checkpoint afirme una procedencia que no tiene.
def rng_state() -> dict[str, Any]: """Snapshot of the Python, NumPy and torch CPU generators.""" name, keys, pos, has_gauss, cached = np.random.get_state(legacy=True) return { "python": json.dumps(random.getstate()), "numpy": json.dumps( [name, np.asarray(keys).tolist(), int(pos), int(has_gauss), float(cached)] ), "torch": torch.get_rng_state(), }
def set_rng_state(state: Mapping[str, Any]) -> None: """Restore a snapshot taken by ``rng_state`` (missing entries are skipped).""" python = state.get("python") if isinstance(python, str): version, keys, gauss = json.loads(python) random.setstate((version, tuple(keys), gauss)) numpy = state.get("numpy") if isinstance(numpy, str): name, keys, pos, has_gauss, cached = json.loads(numpy) np.random.set_state((name, np.array(keys, dtype=np.uint32), pos, has_gauss, cached)) torch_state = state.get("torch") if isinstance(torch_state, torch.Tensor): torch.set_rng_state(torch_state.to(torch.uint8).cpu())El estado de los tres generadores. El de NumPy se guarda desarmado a mano —get_state(legacy=True)
devuelve una tupla con un array de 624 enteros— porque la tupla no es serializable como valores
simples y el requisito de la cabecera es que todo lo sea. El de torch sí es un tensor, que es un
valor permitido.
set_rng_state se salta lo que falte en vez de fallar: un checkpoint viejo sin la entrada de NumPy
se sigue pudiendo cargar.
def save_checkpoint( path: Path, *, step: int, model: nn.Module, optimizer: torch.optim.Optimizer | None, cfg: Mapping[str, Any], model_cfg: Mapping[str, Any], vocab_hash: str | None = None, data_manifest_sha: str | None = None, git_sha: str | None = None, best_val: float | None = None, run_id: str | None = None,) -> Path: """Write one checkpoint atomically (temporary file plus replace) and return its path.
``run_id`` is the MLflow run that wrote it, so ``--resume`` can carry on logging into the same run and the loss curve stays one line instead of two. """ path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) payload: dict[str, Any] = { "step": int(step), "model_state": {k: v.detach().cpu() for k, v in model.state_dict().items()}, "opt_state": optimizer.state_dict() if optimizer is not None else None, "cfg": dict(cfg), "model_cfg": dict(model_cfg), "vocab_hash": vocab_hash, "data_manifest_sha": data_manifest_sha, "git_sha": git_sha, "best_val": best_val, "run_id": run_id, "rng": rng_state(), } tmp = path.with_suffix(path.suffix + ".tmp") torch.save(payload, tmp) tmp.replace(path) return pathEl payload entero. Los tres primeros campos —paso, pesos, estado del optimizador— más el rng son
lo que hace que --resume continúe en vez de reiniciar; los cuatro de procedencia son lo que hace
que alguien pueda decir, seis meses después, con qué partidas y con qué código se entrenó.
Y la escritura es atómica: fichero temporal y replace. Un torch.save directo sobre
best.pt que se interrumpa a mitad —se va la luz, se llena el disco— deja un fichero corrupto donde
antes había uno bueno. Con el temporal, el peor caso es un .tmp huérfano.
def load_checkpoint(path: Path, map_location: str | torch.device = "cpu") -> dict[str, Any]: """Read a checkpoint written by ``save_checkpoint``.""" payload = torch.load(Path(path), map_location=map_location, weights_only=True) if not isinstance(payload, dict) or "model_state" not in payload: raise ValueError(f"{path} is not a rukh checkpoint") return payload
def load_state(model: nn.Module, state: Mapping[str, Any]) -> None: """Load weights, accepting a state dict whose tied head was left out when it was saved.
``MoveDecoder`` ties ``lm_head.weight`` to ``tokens.weight`` in its constructor, so the tied head is already correct once the embedding is loaded: a state dict without it (the one ``rukh publish`` writes, because two names for one tensor is what makes ``safetensors`` refuse the file) loads cleanly and nothing else may be missing. """ missing, unexpected = model.load_state_dict(dict(state), strict=False) absent = [name for name in missing if name != TIED_HEAD] if absent or unexpected: raise ValueError( f"the weights do not match the model: missing {absent}, unexpected {list(unexpected)}" )load_state es el sitio donde vive la consecuencia de los tied embeddings. lm_head.weight y
tokens.weight son el mismo tensor con dos nombres, y lo que rukh publish escribe deja fuera
el nombre atado (dos nombres para un tensor es lo que hace que safetensors rechace el fichero).
Cargar con strict=True fallaría; cargar con strict=False y callar aceptaría cualquier cosa. La
solución es cargar con strict=False y comprobar a mano que lo único que falta es exactamente
el nombre que se sabe que falta.
def load_model( path: Path, map_location: str | torch.device = "cpu") -> tuple[MoveDecoder, dict[str, Any]]: """Rebuild the decoder a checkpoint describes, in eval mode, plus the whole payload.""" payload = load_checkpoint(path, map_location=map_location) model = MoveDecoder(DecoderConfig.model_validate(payload["model_cfg"])) load_state(model, payload["model_state"]) return model.eval(), payload
def restore( payload: Mapping[str, Any], model: nn.Module, optimizer: torch.optim.Optimizer | None = None, with_rng: bool = True,) -> int: """Load weights (and optionally the optimizer and RNG) and return the step reached.""" load_state(model, payload["model_state"]) opt_state = payload.get("opt_state") if optimizer is not None and opt_state is not None: optimizer.load_state_dict(opt_state) rng = payload.get("rng") if with_rng and isinstance(rng, Mapping): set_rng_state(rng) return int(payload["step"])load_model reconstruye el modelo desde el checkpoint, no desde una configuración que le pases:
el model_cfg viaja dentro. Es lo que permite que rukh eval, rukh play, rukh export y
rukh publish model reciban una ruta y ya está. Y devuelve el payload entero además del modelo,
porque los cuatro comandos necesitan algo distinto de él.
Y el paquete, que es la puerta por la que entra todo lo anterior:
"""Training: the loop, the learning-rate schedule and checkpoint handling."""
from rukh.train.checkpoint import ( BEST_NAME, TIED_HEAD, load_checkpoint, load_model, load_state, read_manifest_sha, read_vocab_hash, restore, save_checkpoint, step_name,)from rukh.train.loop import ( TrainConfig, evaluate, param_groups, pick_device, run_dir, skip_batches, train,)from rukh.train.schedule import lr_at
__all__ = [ "BEST_NAME", "TIED_HEAD", "TrainConfig", "evaluate", "load_checkpoint", "load_model", "load_state", "lr_at", "param_groups", "pick_device", "read_manifest_sha", "read_vocab_hash", "restore", "run_dir", "save_checkpoint", "skip_batches", "step_name", "train",]Cuarenta y cinco líneas de reexportes, la tercera lista de este tipo del módulo. rukh.train es lo
que importan el harness de evaluación (load_model, pick_device), el exportador (load_model), el
publicador (load_model, TIED_HEAD) y los labs. Ninguno de ellos sabe que hay tres ficheros
detrás.
"""Tests for rukh.train.checkpoint: round-trip, provenance and RNG restoration."""
from __future__ import annotations
import jsonimport randomfrom pathlib import Path
import numpy as npimport pytestimport torch
from rukh.models import DecoderConfig, MoveDecoderfrom rukh.train import ( TrainConfig, load_checkpoint, read_manifest_sha, read_vocab_hash, restore, save_checkpoint, step_name,)
pytestmark = pytest.mark.unit
TOY = DecoderConfig(vocab_size=32, n_layer=2, n_head=2, d_model=32, block=16)
def make_pair() -> tuple[MoveDecoder, torch.optim.Optimizer]: torch.manual_seed(0) model = MoveDecoder(TOY) return model, torch.optim.AdamW(model.parameters(), lr=1e-3)def test_save_and_load_reproduce_weights_and_step(tmp_path: Path) -> None: model, optimizer = make_pair() idx = torch.randint(1, TOY.vocab_size, (2, 8)) _, loss = model(idx, idx) assert loss is not None loss.backward() optimizer.step()
path = save_checkpoint( tmp_path / step_name(7), step=7, model=model, optimizer=optimizer, cfg=TrainConfig(max_steps=7).model_dump(mode="json"), model_cfg=TOY.model_dump(mode="json"), vocab_hash="abc", data_manifest_sha=None, git_sha="deadbeef", best_val=1.25, ) assert path.name == "step-7.pt" assert not list(tmp_path.glob("*.tmp"))
payload = load_checkpoint(path) assert payload["step"] == 7 assert payload["vocab_hash"] == "abc" assert payload["data_manifest_sha"] is None assert payload["git_sha"] == "deadbeef" assert payload["best_val"] == pytest.approx(1.25) assert payload["cfg"]["max_steps"] == 7 assert payload["model_cfg"]["d_model"] == 32
fresh = MoveDecoder(TOY) fresh_opt = torch.optim.AdamW(fresh.parameters(), lr=1e-3) assert restore(payload, fresh, fresh_opt) == 7 for a, b in zip(model.state_dict().values(), fresh.state_dict().values(), strict=True): assert torch.equal(a, b) with torch.no_grad(): assert torch.allclose(model.eval()(idx)[0], fresh.eval()(idx)[0], atol=1e-6)El test de ida y vuelta da un paso de optimizador antes de guardar, y eso es lo que lo hace
valer: con los pesos recién inicializados, dos modelos con la misma semilla son iguales aunque el
checkpoint no guarde nada. Después compara tensor a tensor y, además, compara los logits, que es la
comprobación que cazaría un estado cargado en el módulo equivocado. Y assert not list(tmp_path.glob("*.tmp")) es el test de la escritura atómica: el temporal no se queda.
def test_restoring_the_rng_repeats_the_same_draws(tmp_path: Path) -> None: model, optimizer = make_pair() random.seed(11) np.random.seed(11) torch.manual_seed(11) path = save_checkpoint( tmp_path / "rng.pt", step=0, model=model, optimizer=optimizer, cfg={}, model_cfg=TOY.model_dump(mode="json"), ) expected = (random.random(), float(np.random.random()), torch.rand(3).tolist())
random.seed(99) np.random.seed(99) torch.manual_seed(99) restore(load_checkpoint(path), MoveDecoder(TOY)) assert (random.random(), float(np.random.random()), torch.rand(3).tolist()) == expectedEl de los generadores es el más bonito del fichero. Guarda, saca tres números de tres bibliotecas distintas, reinicia las tres semillas a otra cosa, restaura el checkpoint y exige los mismos tres números. Es la definición operativa de «reanudar es indistinguible de no haber parado».
def test_vocab_hash_and_manifest_sha_are_defensive(tmp_path: Path) -> None: assert read_vocab_hash(tmp_path) is None (tmp_path / "meta.json").write_text(json.dumps({"n_games": 1}), encoding="utf-8") assert read_vocab_hash(tmp_path) is None (tmp_path / "meta.json").write_text(json.dumps({"vocab_hash": "cafe"}), encoding="utf-8") assert read_vocab_hash(tmp_path) == "cafe"
assert read_manifest_sha(tmp_path / "does-not-exist" / "manifest.json") is None manifest = tmp_path / "manifest.json" manifest.write_bytes(b"{}") sha = read_manifest_sha(manifest) assert sha is not None and len(sha) == 64
def test_a_foreign_file_is_not_a_checkpoint(tmp_path: Path) -> None: path = tmp_path / "other.pt" torch.save({"weights": torch.zeros(2)}, path) with pytest.raises(ValueError, match="not a rukh checkpoint"): load_checkpoint(path)Y los dos últimos: la procedencia que devuelve None en los tres casos en que no la hay (directorio
sin meta.json, meta.json sin la clave, fichero ausente) y un .pt cualquiera que no es un
checkpoint de Rukh. Ese último mensaje de error se lee más veces de las que parece, porque los pesos
de Hugging Face también son un .pt.
tracking.py: que una corrida reanudada sea una sola curva
M0 dejó start_run escribiendo en la base SQLite local. M2 le añade una cosa, y es la que hace que
reanudar no parta la gráfica en dos.
def start_run( name: str, config: Mapping[str, Any], tags: Mapping[str, str] | None = None, run_id: str | None = None,) -> Iterator[Any]: """Start an MLflow run in the local store with ``config`` logged as flattened params.
Tags always include ``rukh_version`` and, when the project lives in a git repo, ``git_sha``. With ``run_id`` the existing run is reopened instead of a new one being created, so a training run that is resumed keeps one curve rather than starting a second one; a run id that cannot be reopened (deleted, or from another store) falls back to a fresh run. Yields the active ``mlflow.ActiveRun``. """El parámetro nuevo es run_id. Va al final y con valor por defecto, así que las llamadas de M0 y M1
siguen valiendo tal cual.
started = None if run_id: try: started = mlflow.start_run(run_id=run_id, tags=run_tags) except Exception as exc: # noqa: BLE001 - a lost run must never stop a training run log.warning("could not reopen the MLflow run %s (%s); starting a new one", run_id, exc) with started or mlflow.start_run(run_name=name, tags=run_tags) as run: params = flatten(config) if params: # A resumed run already carries these; re-logging a changed value is an error there. try: mlflow.log_params(params) except Exception as exc: # noqa: BLE001 - the params are already recorded log.warning("could not log the run parameters: %s", exc) yield runReabrir una corrida existente en vez de crear una nueva. Las dos guardas son la misma idea aplicada
dos veces: un run_id que no se puede reabrir —borrado, o de otra base— cae a una corrida nueva con
un aviso, y volver a registrar los parámetros de una corrida reabierta es un error de MLflow que
tampoco puede costar el entrenamiento. Sin ellas, reanudar una corrida cuyo MLflow se borró tira
cuarenta minutos de GPU por un problema de contabilidad.
Las dos configuraciones
# `small` (12 layers, d=512, 38,971,392 parameters): the course model. Measured at 440k tokens/s# eager on a 5090 (D-024), so the 20,000 steps below are about 40 minutes, not the 4-8 hours the# spec estimated before anyone timed it.preset: smalltokens_dir: data/tokens/uciblock: 200batch_size: 64grad_accum: 4 # 256 effective sequences of 200 tokenslr: 6.0e-4min_lr_ratio: 0.1warmup: 1000max_steps: 20000weight_decay: 0.1betas: [0.9, 0.95]grad_clip: 1.0precision: bf16compile: trueeval_every: 500eval_batches: 50ckpt_every: 1000out_dir: checkpointsseed: 42run_name: small# A timestamp is appended: a second run must not overwrite the step-*.pt series.unique_run_name: trueworkers: 4log_every: 10El modelo del curso. Veintisiete líneas, y ninguna clave es decorativa: BaseConfig rechaza las que
no reconoce, así que lo que está aquí es exactamente lo que TrainConfig sabe leer. Los números son
los de la receta de arriba; el resto son decisiones de operación: eval_every: 500 (cuarenta
evaluaciones en la corrida), ckpt_every: 1000 (veinte checkpoints, que son los que recorre la isla
de «Exportar a ONNX»), workers: 4 para que el cargador no sea el cuello de botella y seed: 42.
El comentario de la cabecera merece leerse: dice 440 000 tokens/s medidos y «unos 40 minutos, no las 4-8 horas que estimaba el spec antes de que nadie lo cronometrara». Dejar la corrección escrita donde estaba la estimación es más útil que borrarla.
# `tiny` (6 layers, d=256, ~5M parameters): the model used to iterate, under 20 minutes on a 5090.preset: tinytokens_dir: data/tokens/uciblock: 200batch_size: 128grad_accum: 2 # 256 effective sequenceslr: 1.0e-3min_lr_ratio: 0.1warmup: 500max_steps: 6000weight_decay: 0.1betas: [0.9, 0.95]grad_clip: 1.0precision: bf16compile: trueeval_every: 250eval_batches: 50ckpt_every: 1000 # every checkpoint feeds the TrainingReplay island of M2out_dir: checkpointsseed: 42run_name: tiny# A timestamp is appended: a second run must not overwrite the step-*.pt series.unique_run_name: trueworkers: 4log_every: 10Y el modelo con el que se itera. Compara las dos: cambian siete claves. El preset; el
batch_size (128 en vez de 64) y el grad_accum (2 en vez de 4), que multiplicados dan las mismas
256 secuencias efectivas, porque un modelo cuatro veces más pequeño cabe entero con lotes del
doble; la tasa de aprendizaje, 1e-3 en vez de 6e-4, porque un modelo pequeño tolera pasos mayores;
max_steps, warmup y eval_every. Fíjate en la proporción del warmup: 500 de 6 000 es el 8 % de
la corrida, frente al 5 % de small. Una corrida corta necesita calentar relativamente más, porque
el riesgo del warmup no es proporcional a los pasos totales sino a lo mal calibrado que está Adam
en los primeros.
La frase del comentario de tiny.yaml —«menos de 20 minutos en una 5090»— también quedó corta: la
tirada real fueron 189 segundos.
Lanzarlo
@app.command("train")def train_cmd( config: Annotated[ Path, typer.Option( "--config", exists=True, dir_okay=False, readable=True, help="Training YAML config." ), ], model_preset: Annotated[ str | None, typer.Option("--preset", help="Override the preset: tiny, small or medium.") ] = None, resume: Annotated[ Path | None, typer.Option( "--resume", exists=True, dir_okay=False, readable=True, help="Checkpoint to continue." ), ] = None, max_steps: Annotated[ int | None, typer.Option("--max-steps", help="Override max_steps from the config.") ] = None,) -> None: """Train a MoveDecoder from a packed token stream, logging the run to MLflow.""" from rukh.config import load_yaml from rukh.models import PRESETS from rukh.train import TrainConfig, train
cfg = load_yaml(config, TrainConfig) if model_preset is not None: if model_preset not in PRESETS: typer.echo(f"error: --preset must be one of {', '.join(PRESETS)}", err=True) raise typer.Exit(code=2) cfg = cfg.model_copy(update={"preset": model_preset, "model": None}) if max_steps is not None: cfg = cfg.model_copy(update={"max_steps": max_steps}) logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s") try: checkpoint = train(cfg, resume=resume) except (FileNotFoundError, ValueError) as exc: typer.echo(f"error: {exc}", err=True) raise typer.Exit(code=1) from exc typer.echo(f"preset: {cfg.preset}") typer.echo(f"steps: {cfg.max_steps}") typer.echo(f"checkpoint: {checkpoint}")El comando, que es lo que la cabecera de cli.py promete: «una cáscara fina sobre la biblioteca».
Carga el YAML, aplica las dos sobreescrituras de la línea de órdenes y llama a train. Las dos
opciones tienen su razón: --preset para lanzar la misma receta con otra talla sin copiar el
fichero, y --max-steps para las pruebas de trescientos pasos del ejercicio de más abajo.
Fíjate en que --preset también pone model: None. Sin eso, una configuración con un model
explícito ignoraría el preset que acabas de pedir por la línea de órdenes y entrenarías otra cosa.
tiny, en tres minutos
uv run rukh train --config configs/train/tiny.yamlLa consola de esta tirada no se capturó. El run sí quedó registrado: estas son las 24 evaluaciones
de validación de tiny-20260919-061533 tal como están en mlruns/mlflow.db, de principio a fin
del entrenamiento (189 segundos en total, poco más de tres minutos):
step 250 val/loss 4.5869 val/top1 0.1903step 500 val/loss 3.5244 val/top1 0.2544step 750 val/loss 3.0194 val/top1 0.2953step 1000 val/loss 2.7751 val/top1 0.3175step 1250 val/loss 2.6150 val/top1 0.3326step 1500 val/loss 2.5098 val/top1 0.3443step 1750 val/loss 2.4252 val/top1 0.3537step 2000 val/loss 2.3687 val/top1 0.3611step 2250 val/loss 2.3210 val/top1 0.3665step 2500 val/loss 2.2764 val/top1 0.3723step 2750 val/loss 2.2394 val/top1 0.3776step 3000 val/loss 2.2115 val/top1 0.3816step 3250 val/loss 2.1809 val/top1 0.3858step 3500 val/loss 2.1552 val/top1 0.3902step 3750 val/loss 2.1312 val/top1 0.3934step 4000 val/loss 2.1111 val/top1 0.3970step 4250 val/loss 2.0932 val/top1 0.3993step 4500 val/loss 2.0770 val/top1 0.4022step 4750 val/loss 2.0622 val/top1 0.4048step 5000 val/loss 2.0500 val/top1 0.4066step 5250 val/loss 2.0398 val/top1 0.4080step 5500 val/loss 2.0316 val/top1 0.4101step 5750 val/loss 2.0251 val/top1 0.4106step 6000 val/loss 2.0200 val/top1 0.4121Mientras corre, abre MLflow (uv run rukh mlflow ui) y mira cuatro series. train/loss debe bajar
deprisa los primeros cientos de pasos y después despacio. val/loss debe acompañarla; cuando se
separen —la de entrenamiento sigue bajando y la de validación se queda plana o sube— has llegado al
sobreajuste y el best.pt ya está guardado del paso anterior. lr debe dibujar la rampa de warmup
y el coseno. Y grad_norm debe estabilizarse: picos recurrentes significan que el clipping está
trabajando demasiado y que la tasa de aprendizaje es alta para este modelo.
Cada 1 000 pasos se escribe un checkpoint, y esos checkpoints son exactamente los que alimentan la repetición del entrenamiento de la lección 7.
small, el modelo del curso
12 capas, 512 dimensiones, 8 cabezas, 38 971 392 parámetros, 20 000 pasos de 64 × 4 secuencias (51 200 tokens por paso, 1 024 millones de tokens en total), tasa 6e-4 con 1 000 de warmup.
Cuánto tardó, ya medido y no estimado: el bucle sostuvo 422 000 tokens/s de mediana en esta
5090 —423 309 en el último paso registrado—, en modo eager (torch.compile no arranca en este
Windows porque no hay Triton, y el bucle cae a eager sin romperse). Mil veinticuatro millones de
tokens a esa velocidad son 2 498 segundos de principio a fin: 42 minutos. El spec del proyecto
estimaba entre cuatro y ocho horas; era una estimación conservadora escrita antes de medir nada, y
la dejamos dicha aquí para que se vea la diferencia entre estimar y medir: el error fue de un
factor de entre seis y once.
uv run rukh train --config configs/train/small.yamlSalida real de la ejecución de referencia (RTX 5090):
2026-09-19 08:56:21,493 step 13000 val/loss 1.5787 val/top1 0.49812026-09-19 08:57:23,733 step 13500 val/loss 1.5695 val/top1 0.50062026-09-19 08:58:25,635 step 14000 val/loss 1.5641 val/top1 0.50192026-09-19 08:59:27,866 step 14500 val/loss 1.5586 val/top1 0.50342026-09-19 09:00:29,794 step 15000 val/loss 1.5528 val/top1 0.50452026-09-19 09:01:32,021 step 15500 val/loss 1.5475 val/top1 0.50602026-09-19 09:02:33,968 step 16000 val/loss 1.5424 val/top1 0.50692026-09-19 09:03:36,099 step 16500 val/loss 1.5378 val/top1 0.50772026-09-19 09:04:37,972 step 17000 val/loss 1.5332 val/top1 0.50872026-09-19 09:05:40,223 step 17500 val/loss 1.5308 val/top1 0.50922026-09-19 09:06:42,081 step 18000 val/loss 1.5275 val/top1 0.51012026-09-19 09:07:44,246 step 18500 val/loss 1.5253 val/top1 0.51102026-09-19 09:08:46,056 step 19000 val/loss 1.5226 val/top1 0.51132026-09-19 09:09:48,207 step 19500 val/loss 1.5201 val/top1 0.51162026-09-19 09:10:50,062 step 20000 val/loss 1.5197 val/top1 0.5120preset: smallsteps: 20000checkpoint: …\rukh\checkpoints\small-20260919-062911\step-20000.ptSesenta y dos segundos entre evaluaciones consecutivas, que son 500 pasos: 8 pasos por segundo, 51 200 tokens cada uno. La cuenta cierra.
Si el proceso se corta —se reinicia la máquina, se cae el driver—, no se pierde nada:
--resume checkpoints/small/step-12000.pt continúa exactamente donde estaba, con el optimizador,
los generadores y la misma corrida de MLflow. Esa es la diferencia entre guardar los pesos y guardar
un checkpoint.
// Ejercicio 01Dos entrenamientos de 300 pasos que explican la receta
Con --max-steps 300, lanza tiny tres veces: (a) tal cual; (b) con warmup: 0 en una copia de
la configuración; (c) con lr: 1.0e-2. Compara las tres curvas de train/loss y de grad_norm
en MLflow. Después responde: ¿cuál de las dos variantes rotas se parece más a la sana en los
primeros diez pasos, y por qué es eso peligroso?
// SoluciónVer la solución
(a) baja suave desde ln(2030) ≈ 7,6. (b) sin warmup pega un salto en los primeros pasos: o la
pérdida sube por encima de 8 antes de volver, o se queda estancada alrededor de 6, que es
aproximadamente la entropía de la distribución de jugadas más frecuentes del corpus; el modelo
ha caído en “predecir siempre lo común”. (c) con lr diez veces mayor, grad_norm se pega al
techo del clipping en casi todos los pasos, señal de que el paso real es el que decide el
clipping y no el optimizador; la pérdida oscila o diverge.
Lo peligroso es la comparación en los primeros diez pasos: las tres bajan. Cualquier
configuración razonable baja de 7,6 a 5 en unos pocos cientos de pasos, simplemente aprendiendo la
frecuencia marginal de cada jugada. Mirar la curva al principio no distingue un entrenamiento sano
de uno roto; hay que esperar a la zona donde la pendiente se suaviza, o mirar métricas que no sean
la pérdida (legalidad, top-1). Es la razón de que eval_every sea 250 en tiny: la señal útil
llega de la validación, no del bucle.
// Ejercicio 02¿Por qué la pérdida de validación va por debajo de la de entrenamiento?
Mira las primeras evaluaciones de tiny y de small: en los primeros checkpoints la pérdida de
validación es menor que la de entrenamiento. Con lo que acabas de leer de loop.py, explica
por qué, y di qué tendría que pasar para que eso fuera un problema.
// SoluciónVer la solución
No es magia ni una fuga al revés: son dos medidas distintas del mismo momento. train/loss se
registra cada log_every pasos y es la media de los micro-lotes de ese paso, que el modelo
vio cuando todavía era peor que ahora; val/loss se mide con los pesos del instante, después de
haber dado esos pasos. Al principio el modelo mejora tan rápido dentro de una ventana de
registro que la diferencia se nota.
Con dropout: 0.0 no hay una segunda explicación posible (con dropout, la de entrenamiento también
sería más alta por tener la red mutilada). Las dos curvas se cruzan en el paso 3 000 de small y
desde ahí el hueco crece de forma monótona, que es lo normal. Sería un problema si la de validación
siguiera por debajo pasados unos miles de pasos: eso significaría que el conjunto de validación es
más fácil que el de entrenamiento —partidas más cortas, jugadores más homogéneos— y que las dos
cifras no se pueden restar.
Qué has aprendido
El bucle entero y las cinco decisiones de la receta, cada una con lo que se rompe si se cambia. Y tres propiedades que no son del modelo sino de la infraestructura, y que valen más que cualquiera de ellas: una corrida reanudada es indistinguible de una que no paró, un checkpoint sabe con qué datos y con qué código se escribió, y perder la traza nunca cuesta una corrida.
Cómo se mide: uv run pytest -m unit -q tests/unit/test_schedule.py tests/unit/test_checkpoint.py
en verde (nueve tests), tiny entrenado en 189 segundos hasta una pérdida de validación de 2,0200
y un top-1 de 41,21 %, y small en 42 minutos hasta 1,5197 y 51,20 %. En MLflow, la
rampa de warmup dibujada y el grad_norm estable.
Lo siguiente es sacarle una jugada: los logits son 2 030 números y convertirlos en e2e4 es una
decisión de diseño con tres perillas y un orden que no es libre.