rukh · lab

// M2 · lección 10

Los labs del decoder

Los cuatro scripts de `labs/m2/` enteros: contar los parámetros contra la fórmula cerrada, demostrar la causalidad con una diferencia de cero exacto, sacar los 96 mapas de atención de una partida y exportar la repetición del entrenamiento checkpoint a checkpoint.

  • tiempo de trabajo135 min
  • además, ejecución sin supervisión+ 10 min de GPU y red
  • nivel base
  • actualizado el22 de septiembre de 2026

Qué vas a construir

Cuatro scripts, 423 líneas, y ninguno entrena nada. Dos son calculadoras que comprueban que entiendes tu propio modelo; los otros dos producen los JSON que dibujan las dos islas de la lección 7.

Están aquí al final porque dependen de todo lo anterior, no porque sean opcionales: los dos primeros son el único punto del módulo en el que una cifra que te ha dado el curso se puede verificar sin ejecutar el modelo, con una calculadora, y esa es la mejor forma que hay de descubrir que tenías una idea equivocada.

Todos los comandos se ejecutan desde la raíz del repositorio rukh.

// Antes de empezarQué cuesta cada lab, y cuál puedes saltarte

LabRutaRelojDejaAtajo
Lab 1Contar los parámetros del decoder contra la fórmula cerradaimprescindiblesegundosel desglose por bloque, que cuadra con el totalno hay
Lab 2Dibujar la máscara causal y demostrar que el futuro no entraimprescindiblesegundosuna diferencia de cero exacto entre la salida completa y la truncadano hay
Lab 3Sacar los 96 mapas de atención de una partida cortasolo observarartifacts/web/m2/attention.json, que dibuja la isla de la lección 7no hay
Lab 4Exportar la repetición del entrenamiento checkpoint a checkpoint desde MLflowsolo observarartifacts/web/m2/training-replay.jsonno hay
imprescindible
Segundos o pocos minutos, y el lab siguiente da por hecho que lo corriste.
solo observar
No se ejecuta nada: se lee la salida y se responde a la pregunta.
labs/m2/README.md
# M2 labs
Scripts referenced by the M2 lesson (`rukh-lab`, `curso/m2/01-el-decoder`). Run them from the
repository root.
| Script | Needs | Lesson lab |
|---|---|---|
| `params.py` | nothing | 1 · contar los parámetros a mano |
| `causal_mask.py` | nothing | 1 · comprobar la causalidad |
| `attention_export.py` | a trained checkpoint | visualización · `artifacts/web/attention.json` |
| `replay_export.py` | an MLflow run and its `step-*.pt` files | visualización · `artifacts/web/training-replay.json` |
`params.py`, `causal_mask.py` and `attention_export.py` are copies of the code blocks in the
lesson: if you edit one, edit the other. `replay_export.py` is only described there, so this file
is its source of truth.
Both JSON outputs are copied into the course with `pnpm sync:data`.

labs/m2/README.mdlíneas 1-17 · p2

Esa frase del final —«params.py, causal_mask.py y attention_export.py son copias de los bloques de código de la lección: si editas uno, edita el otro»— dejó de ser una advertencia y pasó a ser un invariante comprobado: pnpm verify:code compara cada bloque de esta página con el fichero del repositorio en la etiqueta p2 y falla si se separan.

Lab 1 · Contar los parámetros, y comprobar que la cuenta cuadra

Escribir el modelo es la mitad; la otra es saber qué has escrito. Este script recorre el modelo grupo a grupo, suma, y compara con una fórmula escrita a mano. Si las dos cifras no coinciden, tienes una idea equivocada de tu propio modelo.

labs/m2/params.py
"""Parameter count of a MoveDecoder, block by block, against the closed-form formula."""
from rukh.models import MoveDecoder, preset
def formula(cfg) -> int:
"""What the architecture says it should be, computed by hand."""
d, ff, v = cfg.d_model, cfg.ff, cfg.vocab_size
per_block = (
2 * d # ln1
+ (d * 3 * d + 3 * d) # qkv
+ (d * d + d) # attn out projection
+ 2 * d # ln2
+ (d * ff + ff) # mlp in
+ (ff * d + d) # mlp out
)
total = v * d + per_block * cfg.n_layer + 2 * d # tokens + blocks + final ln
if cfg.pos == "learned":
total += cfg.block * d
if not cfg.tie_embeddings:
total += v * d
return total

labs/m2/params.pylíneas 1-22 · p2

La fórmula, y hacer esta cuenta a mano una vez vale más que leer tres artículos. Por bloque: 2·d del primer LayerNorm (ganancia y sesgo por canal), d·3d + 3d de la qkv con su sesgo, d·d + d de la proyección de salida, 2·d del segundo LayerNorm, d·ff + ff de la entrada del MLP y ff·d + d de su salida. Fuera de los bloques: la tabla de tokens (v·d), el LayerNorm final (2·d) y, según la configuración, la tabla de posiciones y una cabeza de salida sin atar.

Los dos condicionales del final son los que hacen que la fórmula compruebe algo: leen cfg.pos y cfg.tie_embeddings, así que la cuenta cambia con la configuración en vez de estar clavada a small.

labs/m2/params.py
for name in ("tiny", "small", "medium"):
cfg = preset(name)
model = MoveDecoder(cfg)
counted = sum(p.numel() for p in model.parameters())
groups = {
"tokens": model.tokens.weight.numel(),
"positions": model.positions.weight.numel() if model.positions is not None else 0,
"blocks": sum(p.numel() for p in model.blocks.parameters()),
"ln_f": sum(p.numel() for p in model.ln_f.parameters()),
"lm_head (tied)": 0 if cfg.tie_embeddings else model.lm_head.weight.numel(),
}
print(f"== {name}: {cfg.n_layer} layers, d={cfg.d_model}, {cfg.n_head} heads ==")
for key, value in groups.items():
print(f" {key:<16} {value:>12,}")
print(f" {'total':<16} {counted:>12,} formula {formula(cfg):>12,}")
print(f" {'non-embedding':<16} {model.num_params():>12,}")
assert counted == formula(cfg), f"{name}: the formula does not match the model"

labs/m2/params.pylíneas 25-41 · p2

El bucle recorre los tres presets y desglosa por grupos. El assert del final es el lab entero: si el modelo construido y la fórmula no coinciden, el script para.

Terminal
uv run python labs/m2/params.py

Salida real de la ejecución de referencia (RTX 5090):

== tiny: 6 layers, d=256, 4 heads ==
tokens 519,680
positions 51,200
blocks 4,738,560
ln_f 512
lm_head (tied) 0
total 5,309,952 formula 5,309,952
non-embedding 5,258,752
== small: 12 layers, d=512, 8 heads ==
tokens 1,039,360
positions 102,400
blocks 37,828,608
ln_f 1,024
lm_head (tied) 0
total 38,971,392 formula 38,971,392
non-embedding 38,868,992
== medium: 16 layers, d=768, 12 heads ==
tokens 1,559,040
positions 153,600
blocks 113,405,952
ln_f 1,536
lm_head (tied) 0
total 115,120,128 formula 115,120,128
non-embedding 114,966,528

Lo que debe salir para small, y que puedes verificar con una calculadora antes de ejecutarlo: embeddings 2 030 × 512 = 1 039 360; posiciones 200 × 512 = 102 400; por bloque 1 024 + 787 968 + 262 656 + 1 024 + 1 050 624 + 1 049 088 = 3 152 384, que por 12 capas son 37 828 608; el LayerNorm final, 1 024. Total 38 971 392, y 38 868 992 sin contar la tabla de posiciones.

Fíjate en el reparto, que es lo que de verdad enseña este lab: el 97 % de los parámetros está en los bloques, y dentro de un bloque, dos tercios están en el MLP. La atención, que es la idea, es la minoría de los pesos.

// Ejercicio 01¿Dónde meterías el próximo millón de parámetros?

Tienes un millón de parámetros extra de presupuesto para small. Calcula, con la fórmula del script, cuánto cuesta cada una de estas tres opciones y cuál cabe: (a) una capa más, (b) subir d_ff de 2 048 a 2 304, (c) subir d_model de 512 a 528 (y ajustar n_head para que siga dividiendo). ¿Cuál elegirías y por qué?

// SoluciónVer la solución

(a) Una capa más son 3 152 384 parámetros: no cabe, se pasa tres veces. (b) Subir d_ff en 256 añade d·256 + 256 + 256·d + d = 262 912 por bloque, 3 154 944 en total: tampoco cabe. (c) Subir d_model a 528 recalcula todo: el término de cada bloque crece un 6,3 %, y sale en torno a 41,4 millones, más de dos millones por encima. Con un millón de margen no cabe ninguna de las tres a lo largo de las doce capas; lo único que cabría es aplicar (b) a la mitad de los bloques, que es una arquitectura rara.

La lección es la que cuenta: en un Transformer el presupuesto no es continuo. Los parámetros van en cuantos de una capa entera o de un incremento de anchura multiplicado por n_layer, y por eso los modelos publicados vienen en tallas y no en cualquier tamaño. Y si hubiera que elegir, el consenso empírico (las leyes de escala de Kaplan y las correcciones de Chinchilla) es que profundidad y anchura deben crecer juntas: un modelo muy profundo y estrecho o muy ancho y plano rinde peor que el cuadrado a igualdad de parámetros. Nuestros tres presets —6×256, 12×512, 16×768— siguen esa diagonal.

Lab 2 · Dibujar la máscara causal y demostrar que funciona

Un test que pasa no te enseña nada si no ves lo que comprueba. Este script dibuja la máscara en la terminal y después hace el experimento que importa: cambiar un token del futuro y comprobar que los logits del pasado no se mueven ni un bit.

labs/m2/causal_mask.py
"""Draw the causal mask and prove causality empirically on a real MoveDecoder."""
import torch
from rukh.models import DecoderConfig, MoveDecoder
T = 8
mask = torch.ones(T, T, dtype=torch.bool).tril()
print("Causal mask (row = query, column = key):")
print(" " + " ".join(f"{j:>2}" for j in range(T)))
for i in range(T):
cells = " ".join(" #" if mask[i, j] else " ." for j in range(T))
print(f" q{i:<3} {cells}")
print(f"\nVisible pairs: {int(mask.sum())} of {T * T} ({100 * mask.float().mean():.1f} %)\n")

labs/m2/causal_mask.pylíneas 1-14 · p2

torch.ones(T, T, dtype=torch.bool).tril() es la máscara causal entera, en una línea: triangular inferior. Es exactamente lo que is_causal=True construye por dentro, y dibujarla con # y . es la diferencia entre creerse una matriz triangular y haberla visto.

labs/m2/causal_mask.py
torch.manual_seed(0)
model = MoveDecoder(DecoderConfig(n_layer=2, n_head=2, d_model=32, vocab_size=64, block=T)).eval()
idx = torch.randint(1, 64, (1, T))
with torch.no_grad():
base, _ = model(idx)
for t in range(T - 1):
changed = idx.clone()
# Replace every token after t with a different one.
changed[:, t + 1 :] = (changed[:, t + 1 :] + 7) % 63 + 1
other, _ = model(changed)
past = (base[:, : t + 1] - other[:, : t + 1]).abs().max().item()
future = (base[:, t + 1 :] - other[:, t + 1 :]).abs().max().item()
print(f" cut after position {t}: max |delta| past {past:.3e} future {future:.3e}")
assert past == 0.0, "the past moved: the mask is not causal"

labs/m2/causal_mask.pylíneas 16-30 · p2

El experimento. Para cada corte, se cambian todos los tokens posteriores y se mide la diferencia máxima en los logits del pasado y del futuro. La aritmética de (changed + 7) % 63 + 1 garantiza que el token nuevo es distinto y sigue estando en el rango válido del vocabulario de juguete.

Terminal
uv run python labs/m2/causal_mask.py

Salida real de la ejecución de referencia (RTX 5090):

Causal mask (row = query, column = key):
0 1 2 3 4 5 6 7
q0 # . . . . . . .
q1 # # . . . . . .
q2 # # # . . . . .
q3 # # # # . . . .
q4 # # # # # . . .
q5 # # # # # # . .
q6 # # # # # # # .
q7 # # # # # # # #
Visible pairs: 36 of 64 (56.2 %)
cut after position 0: max |delta| past 0.000e+00 future 5.934e-01
cut after position 1: max |delta| past 0.000e+00 future 5.969e-01
cut after position 2: max |delta| past 0.000e+00 future 5.884e-01
cut after position 3: max |delta| past 0.000e+00 future 5.925e-01
cut after position 4: max |delta| past 0.000e+00 future 5.863e-01
cut after position 5: max |delta| past 0.000e+00 future 4.818e-01
cut after position 6: max |delta| past 0.000e+00 future 3.524e-01

Lo que tiene que salir: la diferencia en el pasado es exactamente 0,0 —no 1e-7, cero— y la del futuro es grande. Cero exacto porque no es una cuestión de precisión numérica sino de grafo: los tokens futuros ni siquiera entran en el cálculo de las posiciones pasadas. Si alguna vez ves 1e-7 ahí, tienes una fuga sutil, típicamente una normalización aplicada a lo largo del eje de tiempo en vez del de canales.

Las 36 parejas visibles de 64 son T(T+1)/2 sobre : con 8 tokens, el 56 %; con 200, el 50,25 %. La máscara causal cuesta aproximadamente la mitad del cómputo de la atención, y ese es el precio de no hacer trampa.

// Ejercicio 02Rompe la causalidad a propósito

Cambia is_causal=True por is_causal=False en CausalSelfAttention y vuelve a ejecutar el script. Después entrena tiny durante 200 pasos con y sin máscara y compara las dos curvas de pérdida. ¿Cuál baja más rápido? ¿Cuál modelo es mejor? ¿Qué medirías para no dejarte engañar?

Aviso práctico: esto no es una opción de configuración, es una edición de src/rukh/models/decoder.py. Deshazla en cuanto termines (git checkout -- src/rukh/models/), porque un modelo entrenado con ese cambio puesto pasa los tests de forma y falla el de causalidad, y un checkpoint contaminado no se distingue de uno bueno mirando los pesos.

// SoluciónVer la solución

Sin máscara, el assert del script salta en la primera iteración: los logits pasados cambian al tocar el futuro. Entrenando, la curva sin máscara baja muchísimo más rápido y llega a pérdidas que la causal no alcanza jamás, porque la tarea es distinta: con acceso al futuro, predecir el token t+1 se resuelve leyendo el token t+1. En el límite, la pérdida tiende a cero y el top-1 al 100 %.

El modelo es basura y no hay forma de verlo en la curva de entrenamiento. Lo que lo delata es generar: en inferencia el futuro no existe, el modelo recibe una secuencia que no se parece a ninguna que haya visto y propone jugadas al azar. Las dos métricas que lo cazan en un minuto son la legalidad sin máscara (se hunde) y una partida de rukh play (jugadas sin sentido desde el ply 2). Es el caso general de una lección que vale para cualquier proyecto de ML: una pérdida sospechosamente buena es casi siempre una fuga, y la única defensa fiable es una métrica calculada como se va a usar el modelo, no como se entrena.

Lab 3 · Sacar los 96 mapas de atención

Aquí aparece el precio de usar la llamada rápida. F.scaled_dot_product_attention no devuelve la matriz de pesos, así que hay que recalcularla: se engancha un hook a la proyección qkv de cada bloque, se parten sus salidas en Q, K y V, se aplica RoPE si el modelo lo usa y se calcula el softmax a mano con la máscara triangular.

labs/m2/attention_export.py
"""Export the attention weights of one short game to artifacts/web/attention.json."""
import json
from datetime import UTC, datetime
from pathlib import Path
import chess
import torch
from torch.nn import functional as F
from rukh.models.decoder import apply_rope
from rukh.tokenize.uci_vocab import UciTokenizer
from rukh.train import load_model
CKPT = Path("checkpoints/small/best.pt")
OUT = Path("artifacts/web/attention.json")
# Legal's mate: short, famous and every move is easy to follow in the heat map.
MOVES = "e2e4 e7e5 g1f3 b8c6 f1c4 d7d6 b1c3 c8g4 f3e5 g4d1 c4f7 e8e7 c3d5".split()
model, _ = load_model(CKPT)
tok = UciTokenizer()
board = chess.Board()
for uci in MOVES: # fail loudly if the line is not legal
board.push(chess.Move.from_uci(uci))
ids = [tok.bos_id, tok.vocab["<w1800>"], tok.vocab["<b1800>"]]
ids += [tok.vocab[uci] for uci in MOVES]
labels = ["<bos>", "<w1800>", "<b1800>", *MOVES]
idx = torch.tensor([ids], dtype=torch.long)

labs/m2/attention_export.pylíneas 1-29 · p2

La partida es el mate de Légal: trece jugadas, famosa, y cada jugada es fácil de seguir en un mapa de calor. El bucle que la empuja en un tablero antes de nada es una comprobación deliberada —el comentario lo dice: «falla ruidosamente si la línea no es legal»—, porque una línea mal transcrita produciría un mapa perfectamente bonito de una partida que no existe.

Y los labels incluyen los tres tokens de control, que es lo que permite que la isla etiquete las filas y las columnas del <bos> y de los dos Elo. Sin eso, las tres primeras filas del mapa serían misteriosas.

labs/m2/attention_export.py
captured: list[torch.Tensor] = []
def hook(module, args, output): # noqa: ARG001 - torch hook signature
captured.append(output.detach())
handles = [block.attn.qkv.register_forward_hook(hook) for block in model.blocks]
with torch.no_grad():
model(idx)
for handle in handles:
handle.remove()

labs/m2/attention_export.pylíneas 31-42 · p2

Doce hooks, una pasada, y se quitan. Registrar el hook en block.attn.qkv y no en block.attn es lo que da acceso a Q, K y V antes de que la atención los consuma; el hook de attn solo vería la salida ya mezclada.

Quitarlos en el bucle de después no es opcional: un hook que se queda puesto sigue acumulando tensores en captured en cada pasada posterior, y en un proceso largo eso es una fuga de memoria con forma de lista que crece.

labs/m2/attention_export.py
cfg = model.cfg
weights = []
mask = torch.ones(len(ids), len(ids), dtype=torch.bool).tril()
for _layer, qkv in enumerate(captured):
q, k, _ = qkv.split(cfg.d_model, dim=2)
shape = (1, len(ids), cfg.n_head, cfg.head_dim)
q = q.view(shape).transpose(1, 2)
k = k.view(shape).transpose(1, 2)
if cfg.pos == "rope":
q = apply_rope(q, model.rope_cos, model.rope_sin)
k = apply_rope(k, model.rope_cos, model.rope_sin)
scores = (q @ k.transpose(-2, -1)) / (cfg.head_dim**0.5)
probs = F.softmax(scores.masked_fill(~mask, float("-inf")), dim=-1)[0]
weights.append([[[round(v, 4) for v in row] for row in head] for head in probs.tolist()])

labs/m2/attention_export.pylíneas 44-57 · p2

El recálculo, que es CausalSelfAttention.forward sin el paso final. Las mismas formas, el mismo transpose, el mismo RoPE condicional, y después el producto escalar escalado por √head_dim y el softmax con la máscara. probs.tolist() con round(v, 4) recorta a cuatro decimales, que es lo que hace que el JSON pese 211 KB en vez de un megabyte: la isla dibuja color, no necesita el decimal séptimo.

labs/m2/attention_export.py
OUT.parent.mkdir(parents=True, exist_ok=True)
OUT.write_text(
json.dumps(
{
"schema": "rukh-attention/1",
"game": {"moves": labels},
"layers": cfg.n_layer,
"heads": cfg.n_head,
"weights": weights,
"meta": {
"model": "rukh-small",
"checkpoint": str(CKPT),
"generated": datetime.now(UTC).isoformat(timespec="seconds"),
},
}
),
encoding="utf-8",
)
print(f"wrote {OUT} ({cfg.n_layer} layers x {cfg.n_head} heads x {len(ids)}^2)")

labs/m2/attention_export.pylíneas 59-77 · p2

El esquema, con su versión (rukh-attention/1) y su meta.generated. Las dos cosas son requisitos de las islas del curso: el esquema documentado en rukh-lab/src/data/README.md y un marcador de posición que la isla sepa reconocer cuando el JSON está vacío.

Terminal
uv run python labs/m2/attention_export.py

Salida real de la ejecución de referencia (RTX 5090):

wrote artifacts\web\attention.json (12 layers x 8 heads x 16^2)

Una línea y un fichero de 211 KB. El 16^2 es el tamaño de la partida de ejemplo: trece jugadas más los tres tokens de control, 16 × 16 celdas por cabeza, 96 matrices en total. Es el fichero que pnpm sync:data copia a src/data/attention.json y que dibuja AttentionMap en la lección 7.

Lab 4 · La repetición del entrenamiento

El más largo de los cuatro, y el único con opciones de línea de órdenes, porque tiene dos modos que se diferencian en tres órdenes de magnitud de coste.

labs/m2/replay_export.py
"""Export the training curve of a run to ``artifacts/web/training-replay.json``.
Lab 3 of module M2. Reads the metrics MLflow stored during ``rukh train`` and writes **one
entry per checkpoint**: the slider of the ``TrainingReplay`` island walks the ``step-*.pt`` files
the run left behind, not the (much denser) MLflow logging steps, because every other number the
island can show — legality, Elo — only exists for a step whose weights are still on disk.
Only ``step`` is guaranteed in an entry. ``train_loss``, ``val_loss`` and ``val_top1`` are copied
from MLflow when that exact step was logged (with the shipped configs it always is: ``log_every``
and ``eval_every`` both divide ``ckpt_every``), ``legality`` needs ``--with-legality`` and ``elo``
needs ``--with-elo``.
``legality`` is the argmax rate of ``docs/spec/02`` (D-026): the single most likely token, no
temperature, no top-k and no mask. It is the bar ``GOAL.md`` sets at 99 %, and the one that means
something when it is plotted against the training step.
Usage (from the repository root):
uv run python labs/m2/replay_export.py --run-name small-20260919-013000
uv run python labs/m2/replay_export.py --run-id <mlflow run id> --with-legality
"""

labs/m2/replay_export.pylíneas 1-21 · p2

La decisión de diseño está en el primer párrafo: una entrada por checkpoint, no por paso registrado en MLflow. El deslizador de la isla recorre los step-*.pt que dejó la tirada, porque cualquier otra cosa que la isla pueda mostrar —legalidad, Elo— solo existe para un paso cuyos pesos siguen en disco.

Y la segunda: solo step está garantizado. Las tres métricas de MLflow se copian si ese paso exacto se registró (con las configuraciones del repositorio siempre se registra, porque log_every y eval_every dividen a ckpt_every), pero el escritor no lo promete.

labs/m2/replay_export.py
from __future__ import annotations
import argparse
import json
import logging
from collections.abc import Sequence
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
import mlflow
from rukh import paths
from rukh.tracking import tracking_uri
log = logging.getLogger("replay_export")
SCHEMA = "rukh-training-replay/1"
CKPT_GLOB = "step-*.pt"
"""What ``rukh train`` writes every ``ckpt_every`` steps (plus the last one)."""
METRICS = {"train_loss": "train/loss", "val_loss": "val/loss", "val_top1": "val/top1"}
ELO_GAMES = 10
"""Games per rung when ``--with-elo`` is given: enough for a shape, far too few for a number."""

labs/m2/replay_export.pylíneas 23-46 · p2

labs/m2/replay_export.py
def find_run(run_name: str | None, run_id: str | None) -> mlflow.entities.Run:
"""Return the requested run, or the most recent one when nothing is given."""
client = mlflow.tracking.MlflowClient(tracking_uri=tracking_uri(create=False))
if run_id:
return client.get_run(run_id)
experiment = client.get_experiment_by_name("rukh")
if experiment is None:
raise SystemExit("no 'rukh' experiment: run `rukh train` first")
filter_string = f"attributes.run_name = '{run_name}'" if run_name else ""
runs = client.search_runs(
[experiment.experiment_id],
filter_string=filter_string,
order_by=["attributes.start_time DESC"],
max_results=1,
)
if not runs:
raise SystemExit(f"no run matching {run_name or '<latest>'}")
return runs[0]

labs/m2/replay_export.pylíneas 49-66 · p2

Encontrar la ejecución: por id, por nombre, o la más reciente. Los dos raise SystemExit con mensaje en vez de una traza: es un script de línea de órdenes y el usuario necesita saber qué hacer, no dónde petó.

labs/m2/replay_export.py
def history(client: mlflow.tracking.MlflowClient, run_id: str, key: str) -> dict[int, float]:
"""Metric history as ``{step: value}`` (MLflow returns one record per logged point)."""
try:
return {m.step: m.value for m in client.get_metric_history(run_id, key)}
except mlflow.exceptions.MlflowException:
return {}
def run_checkpoints(directory: Path) -> dict[int, Path]:
"""``{step: path}`` of every ``step-*.pt`` in ``directory``, ordered by step."""
found: dict[int, Path] = {}
for path in Path(directory).glob(CKPT_GLOB):
try:
step = int(path.stem.split("-", 1)[1])
except (IndexError, ValueError):
log.debug("ignoring %s: not a step checkpoint", path)
continue
found[step] = path
return dict(sorted(found.items()))

labs/m2/replay_export.pylíneas 69-87 · p2

history devuelve {paso: valor} y se traga la excepción de MLflow devolviendo un diccionario vacío: una métrica que nunca se registró no es un error, es una serie que no existe. Y run_checkpoints ignora los ficheros cuyo nombre no encaja en vez de fallar, porque en el directorio de una ejecución también está best.pt.

labs/m2/replay_export.py
def replay_steps(metric_steps: Sequence[int], checkpoints: dict[int, Path]) -> list[int]:
"""The steps the island scrubs through: the logged steps that kept a checkpoint.
The last logged step is always kept even when its checkpoint is gone (a run that was still
training when this ran, or whose weights were pruned after publication): it is the model the
rest of the table talks about, so leaving it out would be worse than leaving it thin.
"""
chosen = {step for step in metric_steps if step in checkpoints}
if metric_steps:
chosen.add(max(metric_steps))
return sorted(chosen)
def checkpoint_dir(run: mlflow.entities.Run, override: str | None) -> Path:
"""Where the run wrote its checkpoints: ``--checkpoints``, or ``out_dir/<run name>``."""
if override:
return Path(override)
out_dir = str(run.data.params.get("out_dir", "checkpoints"))
return paths.resolve(out_dir) / str(run.info.run_name)

labs/m2/replay_export.pylíneas 90-108 · p2

replay_steps tiene la única regla de negocio interesante del script, y está en su docstring: el último paso registrado se conserva siempre, incluso si su checkpoint ya no está. Es el modelo del que habla el resto de la tabla, así que dejarlo fuera sería peor que dejarlo flaco.

labs/m2/replay_export.py
def elo_of(model: Any, tok: Any, cfg: Any, games: int) -> float | None:
"""Estimated Elo from a handful of games per rung, or None when Stockfish is not around."""
from rukh.engine import EngineNotFound
from rukh.eval.elo import estimate, play_rungs
try:
records = play_rungs(
model,
tok,
cfg.elo_rungs,
games,
cfg.sampling(),
move_time=cfg.elo_move_time,
max_plies=cfg.elo_max_plies,
)
except EngineNotFound as exc:
log.warning("Elo skipped: %s", exc)
return None
if not records:
return None
return estimate(records, samples=cfg.bootstrap, seed=cfg.seed).elo

labs/m2/replay_export.pylíneas 111-131 · p2

labs/m2/replay_export.py
def evaluate_checkpoints(
steps: Sequence[int],
checkpoints: dict[int, Path],
positions: int,
with_elo: bool,
elo_games: int,
) -> dict[int, dict[str, float]]:
"""Measure the legality (and optionally the Elo) of every checkpoint in ``steps``."""
from rukh.eval.legality import legality, sample_positions
from rukh.eval.suite import EvalConfig
from rukh.tokenize.uci_vocab import UciTokenizer
from rukh.train import load_model, pick_device
cfg = EvalConfig(legality_positions=positions, elo_games=elo_games)
games = paths.resolve(cfg.games)
if not games.is_file():
raise SystemExit(f"validation games not found at {games}: --with-legality needs the P1 cut")
tok = UciTokenizer()
prefixes = sample_positions(
games, positions, tok, seed=cfg.seed, pool=cfg.position_pool, block=cfg.block
)
where = cfg.device or pick_device()
measured: dict[int, dict[str, float]] = {}
for step in steps:
ckpt = checkpoints.get(step)
if ckpt is None:
log.warning("step %d has no checkpoint: it keeps its MLflow metrics only", step)
continue
model, _payload = load_model(ckpt, map_location=where)
model = model.to(where).eval()
# `mode="argmax"` on purpose: that is the definition the 99 % bar is written against.
result = legality(model, tok, prefixes, cfg.sampling(), mode="argmax")
values: dict[str, float] = {"legality": round(result.rate, 4)}
line = f" step {step:>6} legality {result.rate:.3f}"
if with_elo:
elo = elo_of(model, tok, cfg, elo_games)
if elo is not None:
values["elo"] = round(elo, 1)
line += f" elo {elo:.0f}"
measured[step] = values
print(line)
return measured

labs/m2/replay_export.pylíneas 134-175 · p2

La parte caza. Por cada checkpoint: cargarlo, medir la legalidad en modo argmax —el comentario lo dice: «a propósito: es la definición contra la que está escrito el listón del 99 %»— y, si se pide, el Elo. Y las posiciones de validación se muestrean una sola vez, fuera del bucle: medir cada checkpoint sobre posiciones distintas haría que la serie no fuera comparable consigo misma, que es exactamente lo único que la serie sirve para mirar.

El import de rukh.eval está dentro de la función y no arriba. Es el mismo patrón que el resto del proyecto: sin --with-legality el script no necesita torch ni el harness, y arranca en un parpadeo.

labs/m2/replay_export.py
def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
)
parser.add_argument("--run-name", default=None, help="MLflow run name (default: latest run)")
parser.add_argument("--run-id", default=None, help="MLflow run id, wins over --run-name")
parser.add_argument(
"--checkpoints",
default=None,
help="directory holding the step-*.pt files (default: <out_dir>/<run name>)",
)
parser.add_argument(
"--out",
default=None,
help="output JSON (default: artifacts/web/training-replay.json)",
)
parser.add_argument(
"--with-legality",
action="store_true",
help="measure the unmasked argmax legality of every checkpoint (slow: one forward pass "
"per sampled position per checkpoint)",
)
parser.add_argument(
"--with-elo",
action="store_true",
help="also estimate the Elo of every checkpoint against Stockfish with --elo-games games "
"per rung. VERY SLOW (eight rungs of real games per checkpoint) and, with so few games, "
"only good for the shape of the curve: the published number comes from `rukh eval`. "
"Implies --with-legality.",
)
parser.add_argument(
"--elo-games",
type=int,
default=ELO_GAMES,
help=f"games per Stockfish rung when --with-elo is given (default: {ELO_GAMES})",
)
parser.add_argument("--positions", type=int, default=500, help="positions per legality check")
return parser.parse_args(argv)

labs/m2/replay_export.pylíneas 178-215 · p2

Los argumentos, con la ayuda de --with-elo escrita como una advertencia en mayúsculas: MUY LENTO (ocho escalones de partidas de verdad por cada checkpoint) y, con tan pocas partidas, solo sirve para la forma de la curva. El número publicado sale de rukh eval. Una opción cuya ayuda explica cuándo no usarla es una opción bien documentada.

labs/m2/replay_export.py
def main() -> None:
logging.basicConfig(level=logging.INFO, format="%(message)s")
args = parse_args()
client = mlflow.tracking.MlflowClient(tracking_uri=tracking_uri(create=False))
run = find_run(args.run_name, args.run_id)
run_id = run.info.run_id
logged = {name: history(client, run_id, key) for name, key in METRICS.items()}
metric_steps = sorted(set().union(*(series.keys() for series in logged.values())))
if not metric_steps:
raise SystemExit(f"run {run_id} has no train/val metrics")
ckpt_dir = checkpoint_dir(run, args.checkpoints)
checkpoints = run_checkpoints(ckpt_dir)
steps = replay_steps(metric_steps, checkpoints)
if not checkpoints:
raise SystemExit(
f"no {CKPT_GLOB} under {ckpt_dir}: the replay is one entry per checkpoint, so point "
"--checkpoints at the directory the run wrote"
)
measured: dict[int, dict[str, float]] = {}
if args.with_legality or args.with_elo:
measured = evaluate_checkpoints(
steps, checkpoints, args.positions, args.with_elo, args.elo_games
)
entries: list[dict[str, float | int]] = []
for step in steps:
entry: dict[str, float | int] = {"step": int(step)}
for name, series in logged.items():
if step in series:
entry[name] = round(series[step], 4)
entry.update(measured.get(step, {}))
entries.append(entry)
default_out = paths.root() / "artifacts" / "web" / "training-replay.json"
out = Path(args.out) if args.out else default_out
out.parent.mkdir(parents=True, exist_ok=True)
payload = {
"schema": SCHEMA,
"steps": entries,
"meta": {
"run": run.info.run_name,
"run_id": run_id,
"preset": run.data.params.get("preset"),
"max_steps": run.data.params.get("max_steps"),
"checkpoints": ckpt_dir.as_posix(),
"generated": datetime.now(UTC).isoformat(timespec="seconds"),
},
}
out.write_text(json.dumps(payload, indent=1) + "\n", encoding="utf-8")
print(f"{len(entries)} checkpoints -> {out}")
if __name__ == "__main__":
main()

labs/m2/replay_export.pylíneas 218-275 · p2

El main, y el raise SystemExit de los checkpoints ausentes con la instrucción dentro («apunta --checkpoints al directorio que escribió la ejecución»). El JSON sale con indent=1 en vez de 2: con veinte entradas de seis campos, un espacio de sangría ahorra un tercio del fichero y se sigue leyendo.

Terminal
# Lo barato: solo las series que MLflow ya registró durante el entrenamiento.
uv run python labs/m2/replay_export.py
# Lo que se exportó de verdad: además, la legalidad sin máscara de cada checkpoint.
uv run python labs/m2/replay_export.py --with-legality --positions 400

Salida real de la segunda ejecución, la de referencia (RTX 5090):

step 1000 legality 0.880
step 2000 legality 0.940
step 3000 legality 0.950
step 4000 legality 0.978
step 5000 legality 0.970
step 6000 legality 0.970
step 7000 legality 0.985
step 8000 legality 0.990
step 9000 legality 0.988
step 10000 legality 0.993
step 11000 legality 0.995
step 12000 legality 0.993
step 13000 legality 0.990
step 14000 legality 0.993
step 15000 legality 0.998
step 16000 legality 0.993
step 17000 legality 0.995
step 18000 legality 0.993
step 19000 legality 0.993
step 20000 legality 0.993
20 checkpoints -> …\rukh\artifacts\web\training-replay.json

Veinte checkpoints, uno cada mil pasos, y una sola de las dos banderas caras. --with-legality cuesta una pasada hacia delante por checkpoint sobre las posiciones que pida --positions —400 aquí, unos segundos por checkpoint— y es lo que llena la fila de legalidad de la isla en los veinte pasos. --with-elo no se usó: juega partidas contra Stockfish en cada checkpoint y convierte un script de un par de minutos en una noche de cómputo, así que el Elo sale como un guion en los veinte.

Y ojo con la resolución de lo que acabas de imprimir: con 400 posiciones, una posición vale un cuarto de punto y el ruido de muestreo ronda el ±1 punto, así que 0.993 y 0.998 son 397 y 399 aciertos de 400 y no dos niveles distintos de competencia. La lectura de esa serie está en la lección 7, debajo de la isla.

labs/m2/README.md
## `replay_export.py`
One entry per **checkpoint**, not per MLflow logging step. The script reads the metric history of
a run, keeps the steps that still have a `step-*.pt` file next to them (plus the last logged step,
whatever happened to its weights) and writes `artifacts/web/training-replay.json` with the schema
`rukh-training-replay/1` documented in `rukh-lab/src/data/README.md`.
Why per checkpoint: the slider of the `TrainingReplay` island is advertised as walking the
checkpoints, and legality and Elo can only be measured for a step whose weights are on disk. With
the shipped configs (`log_every: 10`, `eval_every: 250`/`500`, `ckpt_every: 1000`) every checkpoint
step also has its three MLflow metrics, so the entries come out complete; the writer does not
promise it, though, and only `step` is guaranteed.
```bash
# the curve alone: seconds, needs nothing but the MLflow store
uv run python labs/m2/replay_export.py --run-name small-20260919-013000
# plus the unmasked argmax legality of every checkpoint (the >= 99 % bar of GOAL, D-026)
uv run python labs/m2/replay_export.py --run-name small-... --with-legality --positions 500
# plus a rough Elo per checkpoint: very slow, and 10 games per rung is a shape, not a number
uv run python labs/m2/replay_export.py --run-name small-... --with-elo --elo-games 10
```
| Field | When it is written |
|---|---|
| `step` | always |
| `train_loss`, `val_loss`, `val_top1` | when MLflow logged that metric at that exact step |
| `legality` | with `--with-legality` (or `--with-elo`), unmasked **argmax** rate |
| `elo` | with `--with-elo`, and only when Stockfish is available |
`--checkpoints DIR` overrides where the `step-*.pt` files are looked for; by default it is the
run's `out_dir` parameter joined with the run name, which is where `rukh train` wrote them.
`--with-elo` plays `--elo-games` games against each of the eight Stockfish rungs **for every
checkpoint**, so it is hours of games: it is there to draw the shape of the curve. The published
Elo, with its confidence interval, is the one `rukh eval --suite full` produces.

labs/m2/README.mdlíneas 19-55 · p2

La parte del README que documenta replay_export.py, que es el único de los cuatro que no está copiado de una lección: es más largo que un bloque razonable y su fuente de verdad es su propio README. Bueno: lo era. Ahora está aquí arriba, entero, así que las dos páginas dicen lo mismo y el verificador lo comprueba.

// Ejercicio 03Exporta la atención de una partida tuya

Cambia MOVES de attention_export.py por la partida de rukh play de la lección 4 —los primeros quince plies te sirven— y vuelve a ejecutarlo. Después abre la isla de la lección 7 con tu JSON (pnpm sync:data en rukh-lab) y busca L9H4. ¿Sigue siendo la cabeza de «jugada anterior»?

// SoluciónVer la solución

Debería seguir siéndolo, y ese es el punto del ejercicio: una cabeza especializada es una propiedad de los pesos, no de la partida que le enseñes. Si L9H4 mantiene su media alta sobre la subdiagonal en una partida distinta, la afirmación de la lección 7 gana muchísima fuerza; si se desmorona, lo que se había medido era una casualidad de dieciséis tokens.

Y el otro hallazgo de la lección 7 —los saltos de L5H5 lejos de la diagonal— es el que esperas que no aguante, porque la explicación alternativa era «esa cabeza se ha enganchado a una columna de esta partida concreta». Dos partidas no zanjan la cuestión (harían falta cien y una estadística), pero la segunda partida es barata y ya reparte las probabilidades.

Qué has aprendido

Los cuatro instrumentos del módulo. Dos son calculadoras: una verifica que el modelo que construiste es el que crees, sumando sus parámetros contra una fórmula que escribiste tú, y la otra demuestra la causalidad con una diferencia de cero exacto en vez de confiar en un test que pasa. Los otros dos convierten una tirada de entrenamiento en algo que se puede mirar: 96 matrices de atención y una curva de veinte checkpoints con su legalidad medida.

Cómo se mide: uv run python labs/m2/params.py imprime 38 971 392 para small y su assert pasa; uv run python labs/m2/causal_mask.py imprime 0.000e+00 en las siete líneas del pasado; attention_export.py escribe artifacts/web/attention.json con 12 × 8 × 16² celdas; y replay_export.py --with-legality escribe los veinte checkpoints. Los dos JSON llegan al curso con pnpm sync:data.

Lo siguiente cierra el módulo, y contradice lo que las lecciones anteriores daban por sentado: a small no le faltaba red, le faltaba material.