// 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.
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
| Lab | Ruta | Reloj | Deja | Atajo |
|---|---|---|---|---|
| Lab 1Contar los parámetros del decoder contra la fórmula cerrada | imprescindible | segundos | el desglose por bloque, que cuadra con el total | no hay |
| Lab 2Dibujar la máscara causal y demostrar que el futuro no entra | imprescindible | segundos | una diferencia de cero exacto entre la salida completa y la truncada | no hay |
| Lab 3Sacar los 96 mapas de atención de una partida corta | solo observar | — | artifacts/web/m2/attention.json, que dibuja la isla de la lección 7 | no hay |
| Lab 4Exportar la repetición del entrenamiento checkpoint a checkpoint desde MLflow | solo observar | — | artifacts/web/m2/training-replay.json | no 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.
# M2 labs
Scripts referenced by the M2 lesson (`rukh-lab`, `curso/m2/01-el-decoder`). Run them from therepository 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 thelesson: if you edit one, edit the other. `replay_export.py` is only described there, so this fileis its source of truth.
Both JSON outputs are copied into the course with `pnpm sync:data`.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.
"""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 totalLa 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.
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"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.
uv run python labs/m2/params.pySalida 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,528Lo 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 d² 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.
"""Draw the causal mask and prove causality empirically on a real MoveDecoder."""
import torch
from rukh.models import DecoderConfig, MoveDecoder
T = 8mask = 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")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.
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"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.
uv run python labs/m2/causal_mask.pySalida 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-01Lo 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 T²: 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.
"""Export the attention weights of one short game to artifacts/web/attention.json."""
import jsonfrom datetime import UTC, datetimefrom pathlib import Path
import chessimport torchfrom torch.nn import functional as F
from rukh.models.decoder import apply_ropefrom rukh.tokenize.uci_vocab import UciTokenizerfrom 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)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.
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()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.
cfg = model.cfgweights = []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()])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.
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)")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.
uv run python labs/m2/attention_export.pySalida 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.
"""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 **oneentry per checkpoint**: the slider of the ``TrainingReplay`` island walks the ``step-*.pt`` filesthe run left behind, not the (much denser) MLflow logging steps, because every other number theisland 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 copiedfrom 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, notemperature, no top-k and no mask. It is the bar ``GOAL.md`` sets at 99 %, and the one that meanssomething 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"""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.
from __future__ import annotations
import argparseimport jsonimport loggingfrom collections.abc import Sequencefrom datetime import UTC, datetimefrom pathlib import Pathfrom typing import Any
import mlflow
from rukh import pathsfrom 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."""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]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ó.
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()))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.
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)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.
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).elodef 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 measuredLa 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.
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)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.
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()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.
# 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 400Salida 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.99320 checkpoints -> …\rukh\artifacts\web\training-replay.jsonVeinte 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.
## `replay_export.py`
One entry per **checkpoint**, not per MLflow logging step. The script reads the metric history ofa 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 thecheckpoints, and legality and Elo can only be measured for a step whose weights are on disk. Withthe shipped configs (`log_every: 10`, `eval_every: 250`/`500`, `ckpt_every: 1000`) every checkpointstep also has its three MLflow metrics, so the entries come out complete; the writer does notpromise it, though, and only `step` is guaranteed.
```bash# the curve alone: seconds, needs nothing but the MLflow storeuv 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 numberuv 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 therun'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 everycheckpoint**, so it is hours of games: it is there to draw the shape of the curve. The publishedElo, with its confidence interval, is the one `rukh eval --suite full` produces.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.