// M2 · lección 07
Exportar a ONNX
El puente al navegador: el envoltorio que devuelve solo el último paso, el eje dinámico que se comprueba ejecutando el fichero, las dos cuantizaciones y la paridad de jugada que mide si el `.onnx` elige lo mismo que PyTorch. Con las dos islas que dejan mirar dentro del modelo.
Qué vas a construir
El fichero que descarga el navegador. src/rukh/export/ son 738 líneas que convierten un
checkpoint de PyTorch en tres .onnx —fp32, fp16 e int8— y miden si los tres eligen la misma jugada
que el original. Y al final de la lección, las dos visualizaciones del módulo: la repetición del
entrenamiento checkpoint a checkpoint y el mapa de atención de una partida.
El navegador no ejecuta PyTorch. El puente es ONNXONNXFormato abierto para describir el grafo de una red y sus pesos, independiente del framework que la entrenó. Es lo que permite entrenar en PyTorch y ejecutar en el navegador con `onnxruntime-web`, sin Python ni backend.: un formato que describe
el grafo y los pesos, que onnxruntime-web sabe ejecutar con WebGPU o con WASM.
// Antes de empezarQué cuesta cada lab, y cuál puedes saltarte
| Lab | Ruta | Reloj | Deja | Atajo |
|---|---|---|---|---|
| Exportar a ONNXLos tres ONNX con su comprobación de paridad, y los JSON de las dos islas | cuesta máquina | ~12 min de exportación + minutos de los dos scripts | model{,-fp16,-int8}.onnx, parity.json, attention.json | los tres ONNX están publicados en chorcat/rukh-small |
- cuesta máquina
- Tiempo real de GPU, red o motor. El reloj es el de la RTX 5090 de referencia.
Las dependencias nuevas
M2 añade cinco paquetes al pyproject.toml de M0, y los cinco son de esta lección y de la
siguiente:
"onnx>=1.23", "onnxscript>=0.6", "onnxruntime>=1.30", "onnxconverter-common>=1.16", "safetensors>=0.6",onnx para leer y escribir el fichero, onnxscript porque el exportador moderno de PyTorch lo
necesita, onnxruntime para ejecutar lo exportado (la paridad y la comprobación del eje
dinámico), onnxconverter-common para la conversión a fp16 en condiciones y safetensors para
publicar los pesos en el Hub, que es la lección 9.
Ninguno es un framework de entrenamiento: el decoder sigue escrito a mano.
onnx.py: el grafo
"""Exporting a ``MoveDecoder`` to ONNX for the browser.
The demo only ever needs the distribution over the *next* move, so the exported graph is not themodel itself but a wrapper whose forward returns the last step only: ``(B, V)`` instead of``(B, T, V)``, which divides the output tensor by the context length (200) and saves the browserfrom slicing a megabyte of logits per move.
``torch.onnx.export`` is tried with the dynamo exporter first (the default since torch 2.9 andwhat ``docs/spec/02`` asks for) and falls back to the legacy TorchScript tracer with a warningwhen dynamo is unavailable or fails; which path produced the file is recorded in the metadata,because the two exporters do not emit the same graph and a parity check is only meaningful whenit is known which one ran.
``dynamic_seq`` is not taken on trust. The legacy tracer happily bakes the traced length intothe graph while still being asked for a dynamic axis, and the demo feeds a sequence that growsby one token per move, so the exported file is **run** at two different lengths and the flagreports what actually worked. The model's context (``block``) travels with the file as ONNXmetadata, so the browser knows the limit without being told separately."""Cuatro párrafos y tres decisiones, cada una de las cuales se paga si se omite.
El grafo no es el modelo. La demo solo necesita la distribución de la siguiente jugada, así
que lo que se exporta es un envoltorio cuyo forward devuelve el último paso: (B, V) en vez de
(B, T, V). Eso divide el tensor de salida por la longitud del contexto —doscientas veces menos
números por inferencia— y le ahorra al navegador recortar un megabyte de logits por jugada.
Hay dos exportadores y no dan el mismo grafo. El de dynamo es el moderno y el que pide el diseño; el de TorchScript sigue ahí como respaldo. Cuál corrió se graba en el fichero, porque una prueba de paridad solo significa algo cuando se sabe qué grafo se está comparando.
Un eje dinámico declarado es una promesa, no un hecho. El trazador antiguo hornea alegremente la longitud trazada en el grafo mientras te dice que el eje es dinámico. La demo alarga la secuencia una jugada por turno, así que ese fallo aparecería en el navegador en la segunda jugada de la primera partida.
from __future__ import annotations
import contextlibimport loggingimport sysfrom collections.abc import Iteratorfrom pathlib import Pathfrom typing import Any, Literal
import torchfrom pydantic import BaseModel, ConfigDictfrom torch import Tensor, nn
from rukh import __version__from rukh.models import MoveDecoder
log = logging.getLogger(__name__)
MODEL_NAME = "model.onnx"INPUT_NAME = "idx"OUTPUT_NAME = "logits"BATCH_AXIS = "batch"SEQUENCE_AXIS = "sequence"DEFAULT_OPSET = 18METADATA_PREFIX = "rukh_"class LastStepLogits(nn.Module): """Wraps a decoder so that ``forward(idx)`` returns only the last step's logits."""
def __init__(self, model: MoveDecoder) -> None: super().__init__() self.model = model
def forward(self, idx: Tensor) -> Tensor: logits, _ = self.model(idx) return logits[:, -1, :]Diez líneas, y son la mitad del rendimiento de la demo.
class ExportResult(BaseModel): """What was written and how."""
model_config = ConfigDict(extra="forbid")
path: str exporter: Literal["dynamo", "legacy"] opset: int seq_len: int block: int """The model's context: the exported graph must never be fed more than this many tokens.""" dynamic_batch: bool dynamic_seq: bool """What the file really accepts, not what was asked for: verified by running it.""" dynamic_seq_verified: bool = False """Whether ``dynamic_seq`` was checked by running the file (needs ``onnxruntime``).""" metadata: dict[str, str] = {} """``metadata_props`` written into the file, empty when ``onnx`` is not installed.""" vocab_size: int params: int bytes: int warning: str | None = NoneExportResult tiene dos campos que son la misma idea escrita dos veces: dynamic_seq es lo que
el fichero acepta de verdad, no lo que se pidió, y dynamic_seq_verified dice si se pudo
comprobar (hace falta onnxruntime instalado). Un informe que no distingue «es verdad» de «no lo
he podido mirar» miente por omisión.
def target_path(out: Path) -> Path: """``out`` as a file: a directory gets ``model.onnx`` inside it.""" out = Path(out) return out if out.suffix == ".onnx" else out / MODEL_NAMEdef _dynamic_shapes(dynamic_batch: bool, dynamic_seq: bool) -> dict[str, dict[int, str]] | None: axes: dict[int, str] = {} if dynamic_batch: axes[0] = BATCH_AXIS if dynamic_seq: axes[1] = SEQUENCE_AXIS return {INPUT_NAME: axes} if axes else None
def _dynamic_axes(dynamic_batch: bool, dynamic_seq: bool) -> dict[str, dict[int, str]] | None: shapes = _dynamic_shapes(dynamic_batch, dynamic_seq) if shapes is None: return None axes = dict(shapes) if dynamic_batch: axes[OUTPUT_NAME] = {0: BATCH_AXIS} return axesDos funciones casi iguales porque los dos exportadores esperan formatos distintos: dynamic_shapes
para dynamo, dynamic_axes para el legado. Y el legado necesita además que se le declare el eje de
lote de la salida, que dynamo deduce.
@contextlib.contextmanagerdef _utf8_console() -> Iterator[None]: """Let the exporter print its progress ticks on a legacy Windows code page.
``torch.onnx.export(dynamo=True)`` writes check marks to stdout. On Windows the console is cp1252 by default, so that raises ``UnicodeEncodeError`` inside the exporter and the whole export fails for a reason that has nothing to do with the model. Reconfiguring the streams to replace unencodable characters keeps the modern exporter usable; without it every Windows export silently falls back to the deprecated TorchScript one, which in turn produces an fp16 graph that onnxruntime refuses to load. """ streams = [s for s in (sys.stdout, sys.stderr) if hasattr(s, "reconfigure")] previous = [(s, s.encoding, s.errors) for s in streams] for stream in streams: with contextlib.suppress(Exception): stream.reconfigure(errors="replace") try: yield finally: for stream, encoding, errors in previous: with contextlib.suppress(Exception): stream.reconfigure(encoding=encoding, errors=errors)Veintidós líneas de gestor de contexto para un problema de Windows, y merecen estar aquí porque la
cadena de consecuencias es ejemplar. torch.onnx.export(dynamo=True) imprime marcas de verificación
Unicode en su barra de progreso. La consola de Windows es cp1252 por defecto, así que eso lanza un
UnicodeEncodeError dentro del exportador y la exportación entera falla por un motivo que no
tiene nada que ver con el modelo. El fallo cae al exportador antiguo, que produce un grafo cuyo fp16
onnxruntime se niega a cargar.
Es decir: una configuración regional de la consola acaba en un modelo que el navegador no puede ejecutar. Reconfigurar los flujos para que sustituyan los caracteres que no puede codificar cuesta veinte líneas y rompe la cadena en el primer eslabón.
def export_onnx( ckpt: Path | MoveDecoder, out: Path, opset: int = DEFAULT_OPSET, dynamic_batch: bool = True, seq_len: int = 200, dynamic_seq: bool = True,) -> ExportResult: """Export the next-move head of a checkpoint to ONNX and report which exporter ran.
``ckpt`` may be a checkpoint path or an already-built model (the tests use the latter). ``seq_len`` is the length of the example the exporter traces; with ``dynamic_seq`` the graph still accepts any length up to the model's block, which is what the demo feeds it as a game grows move by move. """ model = ckpt if isinstance(ckpt, MoveDecoder) else _load(Path(ckpt)) model = model.eval() if seq_len > model.cfg.block: raise ValueError(f"seq_len={seq_len} is longer than the model's block {model.cfg.block}") wrapper = LastStepLogits(model).eval() example = torch.zeros((1, seq_len), dtype=torch.long) path = target_path(out) path.parent.mkdir(parents=True, exist_ok=True) common: dict[str, Any] = { "input_names": [INPUT_NAME], "output_names": [OUTPUT_NAME], "opset_version": opset, } warning: str | None = None exporter: Literal["dynamo", "legacy"] = "dynamo" try: with torch.no_grad(), _utf8_console(): torch.onnx.export( wrapper, (example,), str(path), dynamo=True, dynamic_shapes=_dynamic_shapes(dynamic_batch, dynamic_seq), **common, ) except Exception as exc: # noqa: BLE001 - any exporter failure must fall back, not stop warning = f"the dynamo exporter failed ({type(exc).__name__}: {exc}); used the legacy one" log.warning("%s", warning) exporter = "legacy" with torch.no_grad(): torch.onnx.export( wrapper, (example,), str(path), dynamo=False, dynamic_axes=_dynamic_axes(dynamic_batch, dynamic_seq), **common, )El try/except es la caída al exportador antiguo, con el motivo guardado en warning para que
acabe impreso. # noqa: BLE001 con su explicación al lado: cualquier fallo del exportador tiene que
caer, no parar.
works = verify_dynamic_seq(path, seq_len, model.cfg.block) if dynamic_seq and works is False: raise ValueError( f"{path} was exported with a dynamic sequence axis but only runs at length " f"{seq_len}: the {exporter} exporter baked the length in. Re-export with " "dynamic_seq=False and pad the input, or fix the exporter." ) really_dynamic = dynamic_seq if works is None else works metadata = write_metadata( path, { "block": model.cfg.block, "vocab_size": model.cfg.vocab_size, "seq_len": seq_len, "dynamic_batch": dynamic_batch, "dynamic_seq": really_dynamic, "exporter": exporter, "version": __version__, }, ) return ExportResult( path=path.as_posix(), exporter=exporter, opset=opset, seq_len=seq_len, block=model.cfg.block, dynamic_batch=dynamic_batch, dynamic_seq=really_dynamic, dynamic_seq_verified=works is not None, metadata=metadata, vocab_size=model.cfg.vocab_size, params=model.num_params(non_embedding=False), bytes=path.stat().st_size, warning=warning, )Y aquí está la comprobación que da nombre al tercer párrafo del docstring. Si se pidió un eje dinámico y la ejecución a dos longitudes dice que no lo es, se lanza una excepción con instrucciones: reexporta sin eje dinámico y rellena la entrada, o arregla el exportador. No se publica un fichero que va a fallar en el navegador.
really_dynamic = dynamic_seq if works is None else works es la regla de «publica lo que
funcionó»: cuando no se pudo comprobar (sin onnxruntime) se publica lo que se pidió, y
dynamic_seq_verified queda en False para que se sepa.
def sequence_lengths(seq_len: int, block: int) -> list[int]: """Two lengths to run the exported file at: the traced one and a different, legal one.""" other = seq_len + 1 if seq_len < block else max(1, seq_len - 1) return sorted({seq_len, other})def verify_dynamic_seq(path: Path, seq_len: int, block: int) -> bool | None: """Run the file at two sequence lengths; None when ``onnxruntime`` is not installed.""" import numpy as np
try: import onnxruntime as ort except ImportError: log.warning("onnxruntime is not installed: the dynamic sequence axis was not verified") return None lengths = sequence_lengths(seq_len, block) if len(lengths) < 2: # pragma: no cover - block of 1 is not a usable model return None session = ort.InferenceSession(str(path), providers=["CPUExecutionProvider"]) for length in lengths: try: session.run(None, {INPUT_NAME: np.zeros((1, length), dtype=np.int64)}) except Exception as exc: # noqa: BLE001 - any refusal means the axis is not dynamic log.info("the exported graph refused a sequence of %d tokens: %s", length, exc) return False return TrueLas dos longitudes son la trazada y otra distinta y legal: una más si cabe, una menos si no. Es el mínimo para distinguir un eje dinámico de uno horneado, y cuesta dos inferencias.
def set_metadata(path: Path, entries: dict[str, str]) -> dict[str, str]: """Write ``metadata_props`` verbatim into a file; empty when ``onnx`` is not installed.""" try: import onnx except ImportError: log.warning("onnx is not installed: the model metadata (block, vocab) was not written") return {} model = onnx.load(str(path)) onnx.helper.set_model_props(model, entries) onnx.save(model, str(path)) # The dynamo exporter writes the weights next to the graph as `<name>.onnx.data` and points # the initializers at it. `onnx.load` above pulled them into memory and `onnx.save` wrote # them back inline, so the sidecar is now dead weight that would ship to the browser (or be # published to the Hub) without ever being read. Drop it, but only once the file really is # self-contained. sidecar = path.with_suffix(path.suffix + ".data") if sidecar.is_file() and not any( tensor.data_location == onnx.TensorProto.EXTERNAL for tensor in model.graph.initializer ): sidecar.unlink() log.debug("removed the external-data sidecar %s", sidecar.name) return entriesLos metadatos rukh_* viajan dentro del fichero. Un .onnx suelto en un disco tiene que seguir
sabiendo cuál es su ventana de contexto y qué vocabulario habla; guardarlo en un JSON al lado
funciona hasta que alguien mueve el fichero, y alguien lo va a mover.
Y el bloque de las últimas líneas es un hallazgo que solo aparece exportando de verdad: el
exportador de dynamo escribe los pesos en un fichero .onnx.data al lado del grafo y apunta los
inicializadores allí. El onnx.load/onnx.save de aquí arriba los ha traído a memoria y los ha
vuelto a escribir dentro, así que el fichero satélite es peso muerto que se publicaría al Hub o se
enviaría al navegador sin que nadie lo lea. Se borra —pero solo después de comprobar que ningún
inicializador sigue apuntando fuera.
def write_metadata(path: Path, props: dict[str, Any]) -> dict[str, str]: """Write ``rukh_*`` metadata into the file; empty when ``onnx`` is not installed.""" return set_metadata( path, {f"{METADATA_PREFIX}{key}": str(value) for key, value in props.items()} )
def read_metadata(path: Path) -> dict[str, str]: """The ``rukh_*`` metadata of an exported file (needs ``onnx``).""" import onnx
model = onnx.load(str(path)) return {entry.key: entry.value for entry in model.metadata_props}def _load(ckpt: Path) -> MoveDecoder: from rukh.train import load_model
model, _payload = load_model(ckpt) return modelquantize.py: la mitad y la cuarta parte
"""Half precision and dynamic int8 for the two browser backends.
WebGPU runs the fp16 graph (``small`` is about 80 MB that way); the WASM fallback runs the int8one (about 40 MB), which is the difference between a demo that loads on a phone and one thatdoes not. Quantization is restricted to ``MatMul`` and ``Gemm``: those are the weights that makeup almost the whole file, and quantizing the rest costs accuracy for nothing.
``onnxconverter_common`` does the fp16 conversion properly (it keeps the graph's inputs andoutputs in float32 and leaves the ops that overflow in fp32). When it is not installed there isa smaller fallback here that casts the float initializers with ``onnx.numpy_helper`` and puts a``Cast`` back to float32 in front of the output; ``to_fp16`` reports which of the two ran."""WebGPU corre el grafo fp16 (unos 80 MB); el respaldo en WASM corre el int8 (unos 40), que es la
diferencia entre una demo que carga en un móvil y una que no. La cuantizaciónCuantizaciónGuardar y calcular los pesos con menos bits de los que se entrenaron: fp16 (la mitad de tamaño, sin pérdida apreciable) o int8 dinámico (una cuarta parte, con algo de error). En Rukh se cuantizan solo las matrices (`MatMul` y `Gemm`) y cada fichero pasa una prueba de paridad contra PyTorch antes de subirse. se restringe a MatMul y Gemm: son los pesos que ocupan casi todo el
fichero, y cuantizar el resto cuesta exactitud a cambio de nada.
from __future__ import annotations
import loggingfrom pathlib import Pathfrom typing import Literal
from pydantic import BaseModel, ConfigDict
log = logging.getLogger(__name__)
FP16_NAME = "model-fp16.onnx"INT8_NAME = "model-int8.onnx"QUANTIZED_OPS = ["MatMul", "Gemm"]class QuantizeResult(BaseModel): """One converted file."""
model_config = ConfigDict(extra="forbid")
path: str kind: Literal["fp16", "int8"] method: str bytes: int source_bytes: int
@property def ratio(self) -> float: return self.bytes / self.source_bytes if self.source_bytes else 0.0def _sibling(path: Path, name: str, out: Path | None) -> Path: if out is None: return Path(path).with_name(name) out = Path(out) return out if out.suffix == ".onnx" else out / namedef to_fp16(path: Path, out: Path | None = None) -> QuantizeResult: """Convert an fp32 ONNX model to fp16, in place of ``model.onnx``'s sibling by default.""" import onnx
source = Path(path) target = _sibling(source, FP16_NAME, out) target.parent.mkdir(parents=True, exist_ok=True) model = onnx.load(str(source)) try: from onnxconverter_common import float16
converted = float16.convert_float_to_float16(model, keep_io_types=True) method = "onnxconverter_common.float16" except ImportError: converted = _cast_initializers(model) method = "onnx.numpy_helper fallback" log.warning("onnxconverter_common is not installed; used the initializer-casting fallback") onnx.save(converted, str(target)) return QuantizeResult( path=target.as_posix(), kind="fp16", method=method, bytes=target.stat().st_size, source_bytes=source.stat().st_size, )to_fp16 intenta onnxconverter_common y, si no está, cae a un respaldo propio. Y dice cuál
corrió en el campo method, que acaba impreso en la consola. Dos conversiones distintas producen
dos ficheros distintos, así que una paridad del 99,80 % sin saber cuál se usó no es reproducible.
keep_io_types=True es la clave de la conversión buena: las entradas y las salidas del grafo se
quedan en float32 y solo los pesos internos bajan a fp16. El contrato con quien llama al modelo no
cambia.
def _cast_initializers(model: object) -> object: """Fallback fp16 conversion: cast every float initializer and cast the outputs back.
Only the weights change type. Every float-typed value in the graph is retyped to fp16 and a ``Cast`` node is appended so the model still hands the caller float32 logits, which keeps the demo's input and output contract identical to the fp32 file. """ import numpy as np import onnx from onnx import TensorProto, helper, numpy_helper
graph = model.graph # type: ignore[attr-defined] for initializer in graph.initializer: if initializer.data_type != TensorProto.FLOAT: continue array = numpy_helper.to_array(initializer).astype(np.float16) initializer.CopyFrom(numpy_helper.from_array(array, initializer.name)) for value in list(graph.value_info) + list(graph.input): if value.type.tensor_type.elem_type == TensorProto.FLOAT: value.type.tensor_type.elem_type = TensorProto.FLOAT16 for node in graph.node: for attribute in node.attribute: if attribute.type == onnx.AttributeProto.TENSOR: tensor = attribute.t if tensor.data_type == TensorProto.FLOAT: array = numpy_helper.to_array(tensor).astype(np.float16) tensor.CopyFrom(numpy_helper.from_array(array, tensor.name)) for output in graph.output: if output.type.tensor_type.elem_type != TensorProto.FLOAT: continue inner = f"{output.name}_fp16" for node in graph.node: node.output[:] = [inner if name == output.name else name for name in node.output] graph.node.append( helper.make_node("Cast", [inner], [output.name], to=int(TensorProto.FLOAT)) ) return modelEl respaldo, para cuando onnxconverter_common no está. Hace lo mismo a mano —retipa cada
inicializador flotante, retipa los valores intermedios, retipa los tensores de los atributos— y
después añade un nodo Cast delante de cada salida para devolver float32. Esas tres últimas líneas
del bucle de salidas son el truco: se renombra la salida original a <nombre>_fp16, se reescribe
quién la produce y se cuelga el Cast que escribe el nombre público.
def quantize_int8(path: Path, out: Path | None = None) -> QuantizeResult: """Dynamically quantize the MatMul and Gemm weights of an ONNX model to int8.""" from onnxruntime.quantization import QuantType, quantize_dynamic
source = Path(path) target = _sibling(source, INT8_NAME, out) target.parent.mkdir(parents=True, exist_ok=True) quantize_dynamic( model_input=str(source), model_output=str(target), weight_type=QuantType.QInt8, op_types_to_quantize=QUANTIZED_OPS, ) return QuantizeResult( path=target.as_posix(), kind="int8", method="onnxruntime.quantization.quantize_dynamic", bytes=target.stat().st_size, source_bytes=source.stat().st_size, )La cuantización dinámica a int8: 256 niveles y una escala por tensor. El error relativo por peso ronda el 0,4 %, acumulado a lo largo de doce capas, y ya verás lo que eso le hace a las decisiones.
parity.py: ¿elige la misma jugada?
"""Parity between PyTorch and ONNX Runtime: the same move, and how far the logits drifted.
The acceptance criterion of ``docs/spec/02`` is the *move*, not the logits: the demo picks an``argmax`` (or samples from a truncated distribution), so a file that disagrees with PyTorch on0.1 % of positions is fine and one that agrees on the numbers but not on the move is not. Bothare reported: ``agreement`` is the share of positions where the chosen token matches, and``max_abs_logit_delta`` is the worst absolute difference seen anywhere in the logits, which iswhat tells fp32, fp16 and int8 apart.
The positions are the thousand **validation** prefixes of ``docs/spec/02`` §6, the same ones thelegality and accuracy metrics use: a random legal walk visits positions no human would reach, soagreeing on them says little about the file the demo will load. ``random_prefixes`` stays for thecase where no validation parquet is around (and for the tests), and which of the two was used isreported alongside the number."""El criterio es la jugada, no los logits, y el matiz importa: elegir una jugada depende del orden de los logits, no de su valor, así que un error numérico grande puede dejar casi todas las decisiones intactas y uno pequeño puede voltear justo las que estaban ajustadas. Se publican los dos números porque ninguno solo dice lo que pasó.
Y el segundo párrafo es una decisión que se tomó después de la primera versión: las posiciones son las mil de validación, las mismas de la legalidad y la exactitud. Un paseo aleatorio legal visita posiciones a las que ningún humano llega, así que estar de acuerdo en ellas dice poco del fichero que la demo va a cargar. El paseo aleatorio se queda como respaldo, y cuál de los dos se usó se publica al lado del número.
from __future__ import annotations
import loggingimport randomfrom collections.abc import Sequencefrom pathlib import Path
import chessimport numpy as npimport torchfrom pydantic import BaseModel, ConfigDict
from rukh.export.onnx import INPUT_NAME, LastStepLogitsfrom rukh.models import MoveDecoderfrom rukh.tokenize.uci_vocab import UciTokenizer
log = logging.getLogger(__name__)
DEFAULT_N = 1_000MAX_MISMATCHES = 20DEFAULT_GAMES = "data/uci/year=2025/month=02/games.parquet""""The validation month of P1: the parity positions come from here when it exists."""class ParityResult(BaseModel): """How well an exported file reproduces the checkpoint."""
model_config = ConfigDict(extra="forbid")
positions: int agreement: float max_abs_logit_delta: float mismatches: list[int] """Indices of the first few positions where the chosen token differed."""mismatches guarda los índices de las primeras veinte posiciones en las que la jugada cambió. No
entra en ninguna tabla: está ahí para que, cuando una paridad salga mal, se pueda ir a mirar qué
posiciones fallaron en vez de volver a ejecutarlo todo con prints.
def random_prefixes( tok: UciTokenizer, n: int, seed: int = 0, min_ply: int = 1, max_ply: int = 40,) -> list[list[int]]: """Token-id prefixes of random legal games, for a parity check without any dataset.
Every prefix is a real sequence of legal moves, so the model is asked the kind of question the demo asks it, and the walk is seeded so a parity number can be reproduced exactly. """ rng = random.Random(seed) prefixes: list[list[int]] = [] while len(prefixes) < n: board = chess.Board() ids = [ tok.bos_id, tok.vocab["<w1800>"], tok.vocab["<b1800>"], ] plies = rng.randint(min_ply, max_ply) for _ in range(plies): moves = list(board.legal_moves) if not moves: break move = rng.choice(moves) ids.append(tok.vocab.get(move.uci(), tok.unk_id)) board.push(move) prefixes.append(ids) return prefixesdef validation_prefixes( games: Path, tok: UciTokenizer, n: int = DEFAULT_N, block: int = 200, seed: int = 0,) -> list[list[int]]: """Token-id prefixes of ``n`` validation positions, cropped exactly as the model sees them.""" from rukh.eval.legality import sample_positions
positions = sample_positions(games, n, tok, seed=seed, pool=max(n * 5, 10_000), block=block) return [list(position.history) for position in positions]def parity_positions( tok: UciTokenizer, n: int = DEFAULT_N, block: int = 200, seed: int = 0, games: Path | None = None,) -> tuple[list[list[int]], str, str | None]: """``(prefixes, source, warning)``: validation positions, or random walks with a warning.""" from rukh import paths
parquet = Path(games) if games is not None else paths.resolve(DEFAULT_GAMES) if parquet.is_file(): prefixes = validation_prefixes(parquet, tok, n=n, block=block, seed=seed) if prefixes: return prefixes, "validation", None warning = ( f"no validation games at {parquet}: parity was measured on random legal walks, which " "visit positions no human would reach" ) log.warning("%s", warning) return random_prefixes(tok, n, seed=seed), "random-walk", warningLas tres fuentes de posiciones, con la de verdad primero y el aviso escrito por si se cae al paseo
aleatorio. Fíjate en que validation_prefixes reutiliza sample_positions de eval.legality: la
paridad se mide sobre exactamente los mismos prefijos que la legalidad, con el mismo recorte y la
misma semilla.
def _session(onnx_path: Path): # type: ignore[no-untyped-def] """An ONNX Runtime session on the CPU provider (tests monkeypatch this).""" import onnxruntime as ort
return ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"])def parity( ckpt: Path | MoveDecoder, onnx_path: Path, positions: Sequence[Sequence[int]], n: int = DEFAULT_N,) -> ParityResult: """Compare the argmax token and the logits of the checkpoint and the exported file.""" model = ckpt if isinstance(ckpt, MoveDecoder) else _load(Path(ckpt)) wrapper = LastStepLogits(model.eval()).eval() session = _session(Path(onnx_path)) used = list(positions)[:n] if not used: raise ValueError("parity needs at least one position")
agreed = 0 worst = 0.0 mismatches: list[int] = [] for index, history in enumerate(used): idx = np.asarray([list(history)], dtype=np.int64) with torch.no_grad(): reference = wrapper(torch.from_numpy(idx)).numpy()[0].astype(np.float64) exported = np.asarray(session.run(None, {INPUT_NAME: idx})[0])[0].astype(np.float64) worst = max(worst, float(np.max(np.abs(reference - exported)))) if int(np.argmax(reference)) == int(np.argmax(exported)): agreed += 1 elif len(mismatches) < MAX_MISMATCHES: mismatches.append(index) return ParityResult( positions=len(used), agreement=agreed / len(used), max_abs_logit_delta=worst, mismatches=mismatches, )La comparación. PyTorch y ONNX reciben el mismo array de enteros, y las dos salidas se suben a
float64 antes de restarlas para que la diferencia máxima no dependa de la precisión en la que
casualmente venía cada una. El acuerdo se cuenta comparando los dos argmax; la diferencia máxima
se acumula sobre todos los logits de todas las posiciones.
Y el envoltorio del lado de PyTorch es LastStepLogits, el mismo que se exportó. Comparar el
modelo completo contra el grafo del último paso compararía dos cosas distintas.
def _load(ckpt: Path) -> MoveDecoder: from rukh.train import load_model
model, _payload = load_model(ckpt) return model__init__.py: rukh export en una función
"""Export: ONNX for the browser, fp16 and int8 quantization and PyTorch parity.
``export_all`` is the whole of ``rukh export``: write the fp32 graph, derive the fp16 and int8files the two browser backends need, and check every file it produced against PyTorch."""from __future__ import annotations
from pathlib import Path
from pydantic import BaseModel, ConfigDict
from rukh.export.onnx import ( DEFAULT_OPSET, INPUT_NAME, MODEL_NAME, OUTPUT_NAME, ExportResult, LastStepLogits, export_onnx, read_metadata, sequence_lengths, set_metadata, target_path, verify_dynamic_seq, write_metadata,)from rukh.export.parity import ( DEFAULT_GAMES, ParityResult, parity, parity_positions, random_prefixes, validation_prefixes,)from rukh.export.quantize import ( FP16_NAME, INT8_NAME, QUANTIZED_OPS, QuantizeResult, quantize_int8, to_fp16,)class ExportBundle(BaseModel): """Everything one ``rukh export`` produced."""
model_config = ConfigDict(extra="forbid")
onnx: ExportResult fp16: QuantizeResult | None = None int8: QuantizeResult | None = None parity: dict[str, ParityResult] = {} """Parity per file kind: ``fp32``, ``fp16``, ``int8``.""" parity_source: str | None = None """``validation`` (the positions of ``docs/spec/02`` §6) or ``random-walk``.""" parity_warning: str | None = Nonedef export_all( ckpt: Path, out: Path, opset: int = DEFAULT_OPSET, seq_len: int = 200, fp16: bool = False, int8: bool = False, check_parity: bool = False, positions: int = 1_000, seed: int = 0, games: Path | None = None,) -> ExportBundle: """Export, quantize and check one checkpoint in a single pass.
Parity is measured on validation positions (``games``, or the validation month of P1 when it is on disk) and falls back to random legal walks with a warning when there are none. """ from rukh.tokenize.uci_vocab import UciTokenizer from rukh.train import load_model
model, _payload = load_model(Path(ckpt)) bundle = ExportBundle(onnx=export_onnx(model, out, opset=opset, seq_len=seq_len)) if fp16: bundle.fp16 = to_fp16(Path(bundle.onnx.path)) if int8: bundle.int8 = quantize_int8(Path(bundle.onnx.path)) # The demo loads the fp16 or the int8 file, not the fp32 one, so the context length has to # travel with them too; neither converter promises to keep the metadata of its input. for quantized in (bundle.fp16, bundle.int8): if quantized is not None and bundle.onnx.metadata: set_metadata(Path(quantized.path), bundle.onnx.metadata) if check_parity: prefixes, source, warning = parity_positions( UciTokenizer(), n=positions, block=model.cfg.block, seed=seed, games=games ) bundle.parity_source = source bundle.parity_warning = warning checks = {"fp32": bundle.onnx.path} if bundle.fp16 is not None: checks["fp16"] = bundle.fp16.path if bundle.int8 is not None: checks["int8"] = bundle.int8.path bundle.parity = { kind: parity(model, Path(path), prefixes, n=positions) for kind, path in checks.items() } return bundleexport_all es el comando entero: exportar, cuantizar y comprobar todo lo que ha producido. Dos
detalles que solo se descubren exportando de verdad:
- Los metadatos se vuelven a escribir en los ficheros cuantizados. El comentario lo dice: la
demo carga el fp16 o el int8, no el fp32, y ninguno de los dos conversores promete conservar los
metadata_propsde su entrada. - Las posiciones de paridad se sacan una vez y se usan para los tres ficheros. Comparar fp16 e int8 sobre muestras distintas haría que sus dos cifras no se pudieran poner en la misma tabla.
__all__ = [ "DEFAULT_GAMES", "DEFAULT_OPSET", "FP16_NAME", "INPUT_NAME", "INT8_NAME", "MODEL_NAME", "OUTPUT_NAME", "QUANTIZED_OPS", "ExportBundle", "ExportResult", "LastStepLogits", "ParityResult", "QuantizeResult", "export_all", "export_onnx", "parity", "parity_positions", "quantize_int8", "random_prefixes", "read_metadata", "sequence_lengths", "set_metadata", "target_path", "to_fp16", "validation_prefixes", "verify_dynamic_seq", "write_metadata",]El comando, y lo que midió
@app.command("export")def export_cmd( ckpt: Annotated[ Path, typer.Option( "--ckpt", exists=True, dir_okay=False, readable=True, help="Checkpoint to export." ), ], out: Annotated[Path, typer.Option("--out", help="Output directory (or .onnx file).")], opset: Annotated[int, typer.Option("--opset", help="ONNX opset version.")] = 18, seq_len: Annotated[ int, typer.Option("--seq-len", help="Example length the exporter traces.") ] = 200, fp16: Annotated[bool, typer.Option("--fp16", help="Also write the fp16 model.")] = False, int8: Annotated[bool, typer.Option("--int8", help="Also write the int8 model.")] = False, check_parity: Annotated[ bool, typer.Option("--check-parity", help="Compare every file with PyTorch.") ] = False, positions: Annotated[ int, typer.Option("--positions", help="Positions used by the parity check.") ] = 1000, games: Annotated[ Path | None, typer.Option( "--games", exists=True, dir_okay=False, readable=True, help="Validation games parquet for the parity positions.", ), ] = None, as_json: Annotated[bool, typer.Option("--json", help="Print the result as JSON only.")] = False,) -> None: """Export the next-move head to ONNX, quantize it and check parity with PyTorch.""" from rukh.export import export_all
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s") try: bundle = export_all( ckpt, out, opset=opset, seq_len=seq_len, fp16=fp16, int8=int8, check_parity=check_parity, positions=positions, games=games, ) except (ImportError, ValueError) as exc: typer.echo(f"error: {exc}", err=True) raise typer.Exit(code=1) from exc if as_json: typer.echo(bundle.model_dump_json(indent=2)) return typer.echo(f"exporter: {bundle.onnx.exporter} (opset {bundle.onnx.opset})") if bundle.onnx.warning: typer.echo(f"warning: {bundle.onnx.warning}") checked = "verified" if bundle.onnx.dynamic_seq_verified else "not verified" typer.echo( f"shapes: batch dynamic={bundle.onnx.dynamic_batch}, " f"sequence dynamic={bundle.onnx.dynamic_seq} ({checked}), block={bundle.onnx.block}" ) typer.echo(f"fp32: {bundle.onnx.path} ({bundle.onnx.bytes} bytes)") for quantized in (bundle.fp16, bundle.int8): if quantized is not None: typer.echo( f"{quantized.kind}: {quantized.path} ({quantized.bytes} bytes, " f"{quantized.ratio:.2f} of fp32, {quantized.method})" ) for kind, result in bundle.parity.items(): typer.echo( f"parity {kind}: {result.agreement:.4f} on {result.positions} " f"{bundle.parity_source} positions " f"(max |delta logits| {result.max_abs_logit_delta:.4g})" ) if bundle.parity_warning: typer.echo(f"warning: {bundle.parity_warning}")uv run rukh export --ckpt checkpoints/small/best.pt --out artifacts/onnx/small --fp16 --int8 --check-paritySalida real de la ejecución de referencia (RTX 5090):
exporter: dynamo (opset 18)kind: decoder -> logitsshapes: batch dynamic=True, sequence dynamic=True (verified), block=200fp32: artifacts/onnx/small/model.onnx (156743532 bytes)fp16: artifacts/onnx/small/model-fp16.onnx (78804476 bytes, 0.50 of fp32, onnxconverter_common.float16)int8: artifacts/onnx/small/model-int8.onnx (43459142 bytes, 0.28 of fp32, onnxruntime.quantization.quantize_dynamic)parity fp32: 1.0000 on 1000 validation positions (max |delta logits| 2.551e-05)parity fp16: 0.9980 on 1000 validation positions (max |delta logits| 0.01404)parity int8: 0.9540 on 1000 validation positions (max |delta logits| 1.75)Salió por el exportador dynamo, el eje de secuencia está verificado —no declarado— y las posiciones fueron las de validación, no un paseo aleatorio. Las tres cosas se leen en la salida porque las tres podían haber salido de otra manera.
Y aquí está la segunda mala noticia del módulo: el listón del 99,9 % no lo pasa ninguna de las dos cuantizaciones.
| Fichero | Tamaño | Paridad de jugada | Max |Δ logits| | ¿Pasa el ≥ 99,9 %? |
|---|---|---|---|---|
| fp32 | 156,7 MB | 100 % | 2,551e-05 | sí |
| fp16 | 78,8 MB | 99,80 % | 0,01404 | no |
| int8 | 43,5 MB | 95,40 % | 1,75 | no |
fp16 se queda a una décima de punto del listón: elige otra jugada en 2 de cada 1 000 posiciones cuando el 99,9 % admitía una, y sus logits se mueven una centésima. Eso es prácticamente el mismo modelo, y como se sirve por WebGPU es el que juega en la demo cuando el navegador puede.
int8 es otra cosa. Cambia la jugada en 46 de cada 1 000 posiciones y mueve los logits hasta 1,75, que es un margen del tamaño de la diferencia habitual entre la primera y la segunda jugada. Un 95,40 % de paridad no es «casi el mismo modelo»: es un modelo medible y distinto, un 4,6 % de sus decisiones cambiadas, y no se ha medido cuánto Elo cuesta esa diferencia. Y es precisamente el fichero que descarga un móvil, porque es el respaldo en WASM cuando no hay WebGPU: los 43,5 MB frente a los 78,8 del fp16 son la razón de que exista. Así que quien juegue desde un teléfono antiguo está jugando contra una variante del modelo que la tabla única no mide.
Es la misma foto guardada en JPEG a dos calidades. A calidad 95 el fichero pesa la mitad y nadie distingue la copia del original salvo con una lupa sobre un borde: eso es fp16. A calidad 40 pesa un cuarto, de lejos sigue siendo la misma foto y de cerca los bloques se ven; y si lo que había que leer en ella era un texto pequeño —una posición con dos jugadas casi empatadas—, ahí es donde se pierde: eso es int8. Cuánto se pierde no se sabe mirando el tamaño del fichero; hay que mirar la foto, y aquí, jugar las partidas.
Lo honesto es dejarlo escrito así y no redondear: el listón del hito era ≥ 99,9 % de paridad de jugada para las tres variantes, fp32 lo cumple, fp16 se queda a 99,80 % e int8 a 95,40 %. Medir el Elo del int8 con su propia tirada de partidas es trabajo para un hito posterior.
// Ejercicio 01¿Por qué int8 cambia los logits y no la jugada?
Mira la salida de la exportación: la diferencia máxima de logits del fichero int8 es más de cien veces la del fp16 (1,75 frente a 0,01404), y sin embargo su paridad de jugada sigue siendo del 95,40 %, no del 5 %. Explica por qué el desastre en los logits se traduce en un desperfecto moderado en las jugadas, y di en qué tipo de posición esperarías que int8 falle.
// SoluciónVer la solución
La cuantización dinámica a int8 representa cada matriz con 256 niveles y una escala por tensor: el error relativo por peso es del orden del 0,4 %, y ese error se acumula a lo largo de doce capas. Los logits resultantes se desplazan de forma perceptible. Pero elegir una jugada no depende del valor de los logits sino de su orden, y en la mayoría de posiciones la diferencia entre el primero y el segundo es mucho mayor que el ruido introducido. fp16, con 10 bits de mantisa y el mismo rango relativo en toda la escala, introduce un error unas cien veces menor.
Int8 falla justo donde el margen es pequeño: posiciones en las que dos o tres jugadas están casi empatadas —aperturas muy conocidas con varias continuaciones estándar, finales simples donde varias jugadas mantienen la evaluación— y también en las posiciones donde el modelo está poco seguro en general. Ahí las dos jugadas suelen ser igual de razonables, así que el efecto sobre la fuerza es probablemente menor de lo que sugiere un 4,6 % de decisiones cambiadas. Pero fíjate en el «probablemente»: nadie lo ha medido, y esa discrepancia no se reparte uniformemente. Por eso la paridad se mide sobre posiciones de partidas reales y no sobre ruido, y por eso el 95,40 % se publica como un incumplimiento del listón y no como una nota al pie.
Ver el entrenamiento: TrainingReplay
Al terminar esta sección sabrás leer una curva de entrenamiento como quien lee un electrocardiograma: dónde está sano, dónde se está sobreajustando y en qué momento el modelo dejó de decir tonterías.
La isla de abajo recorre los checkpoints de la tirada real: un paso del deslizador, un
step-*.pt. El deslizador elige uno; el gráfico dibuja la pérdida de entrenamiento y la de
validación contra el eje izquierdo y el top-1 de siguiente jugada contra el derecho, con una línea
vertical en el checkpoint seleccionado, y debajo están sus números exactos, incluidas la legalidad
sin máscara por argmax —medida en los veinte checkpoints— y el Elo estimado, que en esta tirada no
se midió en ninguno. Los campos que no se midieron salen como un guion, y su línea deja un hueco en
vez de bajar a cero. El script que produce ese JSON es labs/m2/replay_export.py, en la lección 10.
Tres cosas concretas que mirar, en este orden:
- Dónde se separan las dos pérdidas. Mientras van juntas, el modelo generaliza: lo que aprende
del entrenamiento le sirve para partidas que no ha visto. El punto donde la de validación se
aplana mientras la otra sigue bajando es el principio del sobreajuste, y es el punto que decide
cuál es el
best.pt. - Cuánto compra la segunda mitad del entrenamiento. Lleva el deslizador al paso 10 000 y compara con el 20 000. La diferencia entre esas dos cifras de top-1 es, literalmente, lo que valieron veinte minutos de GPU, y es el número con el que se decide si la tirada siguiente dura más pasos o cambia otra cosa.
- Cuándo despega la legalidad. Mira solo esa fila mientras llevas el deslizador del paso 1 000 al 8 000: 88,0 % → 99,0 %. Y luego sigue hasta el final y comprueba que ya no pasa nada.
pérdida de entrenamiento pérdida de validación top-1 de validación legalidad sin máscara (barras)
- Paso
- 20000
- Pérdida (train)
- 1,459
- Pérdida (val)
- 1,520
- Top-1 (val)
- 51,2 %
- Legalidad sin máscara
- 99,3 %
- Elo estimado
- —
Lo que dibuja la isla, leído de los números y no de la forma de la curva. Los 20 checkpoints van de 2,6531 de pérdida de entrenamiento y 2,6279 de validación en el paso 1 000 a 1,4594 y 1,5197 en el 20 000, con el top-1 de validación subiendo de 0,3341 a 0,5120. Casi todo el descenso ocurre pronto: del paso 1 000 al 5 000 la validación cae 0,85 (2,6279 → 1,7822) y en los 15 000 pasos restantes solo 0,26 más. El top-1 hace exactamente lo mismo: +12,2 puntos en los primeros 5 000 pasos y +5,6 en los 15 000 siguientes; la segunda mitad entera del entrenamiento, del paso 10 000 al 20 000, compró 2,4 puntos de top-1.
El hueco entre las dos pérdidas responde a la pregunta del sobreajuste, y la respuesta no es la que
uno espera. En los pasos 1 000 y 2 000 la validación va por debajo de la de entrenamiento
(−0,025 y −0,015), por el motivo que fijaba el segundo ejercicio de la lección 3. Las dos curvas se
cruzan en el paso 3 000 y desde ahí el hueco crece de forma monótona: +0,021 en el 10 000,
+0,046 en el 15 000 y +0,060 en el 20 000. Sesenta milésimas sobre una pérdida de 1,52 es un
4 %. Hay sobreajuste, sí, y se ve; pero la curva de validación no se ha aplanado: en los
últimos mil pasos todavía baja 0,003 y el top-1 todavía sube. Con 240 millones de tokens de
entrenamiento —poco más de cuatro épocas de enero— 39 millones de parámetros no son demasiados
parámetros. Por eso el best.pt de esta tirada es el último checkpoint y no uno intermedio
rescatado antes del desastre.
La conclusión del módulo al cerrarse fue que a small no lo limitaba el sobreajuste sino el
presupuesto, y que entrenar más pasos era la mejora barata que faltaba por probar. Se probó
después, y la respuesta fue otra: ese hueco que crece mientras el top-1 se aplana es la firma de
repetir material, no la de sobrar parámetros. La última lección del módulo tiene la historia
entera y la tabla que lo demuestra.
Y ahora la fila que hace de este módulo el de ver emerger el juego. El modelo no tiene tablero, no tiene reglas y no sabe qué es un alfil: solo ve secuencias de cuatro caracteres. En el paso 1 000 propone una jugada imposible en el 12 % de las posiciones —legalidad 88,0 %, 352 aciertos de 400—; en el paso 8 000 ya respeta las reglas en el 99,0 %, 396 de 400. Once puntos en siete mil pasos. Y ahí se para: del paso 8 000 al 20 000 la legalidad pasa de 99,0 % a 99,25 %, un cuarto de punto —una posición de 400— en los 12 000 pasos restantes, mientras la pérdida de validación todavía baja de 1,6727 a 1,5197 y el top-1 sube de 47,78 % a 51,20 %. La legalidad satura mucho antes que la pérdida.
Antes de sacarle punta a la curva, el ruido. Cada punto se midió sobre 400 posiciones de validación, así que un acierto o un fallo valen un cuarto de punto y el error de muestreo ronda ±1 punto. Por eso a partir del paso 10 000 la serie no sube en línea recta sino que vibra entre 99,0 % y 99,75 %: el 99,75 % del paso 15 000 y el 99,25 % del 20 000 son 399 y 397 posiciones legales de 400, una diferencia de dos posiciones. Quien concluya de ahí que el checkpoint 15 000 sabe más reglas que el final está leyendo ruido como si fuera señal. Lo único que esa parte de la curva dice es «ya está»; lo que significa algo es el tramo de subida, del paso 1 000 al 8 000.
Y de ahí sale el argumento más importante de la lección: legalidad y fuerza son cosas distintas. En el paso 8 000 el modelo ya cumple las reglas casi perfectamente y juega bastante peor que en el 20 000: misma legalidad dentro del ruido, pero una pérdida de validación un 10 % más alta y 3,4 puntos menos de top-1, medidos sobre las mismas partidas de validación. Saber que un alfil va en diagonal no es saber adónde conviene moverlo. Por eso el listón del ≥ 99 % del hito es un requisito de entrada y no una medida de calidad, y por eso la fila de la tabla única no se queda en la legalidad y trae además top-1, puzles, Δcp y Elo.
Lo que sigue abierto es la otra pregunta, la cara: el Elo por checkpoint. La isla tiene sitio
para él y en esta tirada sale como un guion en los veinte, porque --with-elo juega partidas
contra Stockfish en cada uno y cuesta una noche de cómputo. Sin ese dato se puede decir cuánta
imitación compró cada tramo de pasos, pero no cuánta fuerza; y con la legalidad ya saturada desde
el paso 8 000, la fuerza es justo lo que falta por ver. Queda anotada como tal, no resuelta de
memoria.
Ver la atención: AttentionMap
Al terminar esta sección sabrás mirar una matriz de atención sin contarte una película: qué patrones son reales, cuáles son artefactos y qué no se puede concluir de un mapa bonito.
Abajo está la matriz de una cabeza de una capa sobre una partida corta (el mate de Légal, trece jugadas más los tokens de control). Cada fila es una jugada que mira; cada columna, una jugada mirada. La mitad superior derecha está siempre vacía: eso es la máscara causal, que ahora puedes ver en lugar de creértela. El color es relativo al máximo de la cabeza seleccionada, así que comparar intensidades entre cabezas distintas no significa nada; comparar dentro de una, sí. El valor exacto de cada celda está en su título y en el resumen de arriba, que te dice las tres jugadas a las que más mira la fila que elijas con el selector.
Qué buscar mientras cambias de capa y de cabeza:
- La subdiagonal encendida. La casilla inmediatamente a la izquierda de la diagonal es «la jugada anterior»; la diagonal misma es cada jugada mirándose a sí. Suele haber alguna cabeza dedicada a la anterior, y tiene todo el sentido: es el vecino más informativo de una secuencia. En este modelo la hay —una sola, y no en la capa en la que la buscarías—; la isla abre en ella (capa 9, cabeza 4) y la tabla que sigue dice con qué números.
- Columnas verticales. Una columna entera encendida es una jugada a la que mira todo el mundo,
normalmente porque acumula contexto. El
<bos>del principio suele comportarse así en varias cabezas y, en los Transformers de texto, se conoce como sumidero de atención: la cabeza no tiene nada que mirar y vuelca su peso en el primer token, que es su forma de decir «no aporto nada aquí». Es lo que hace alguien en una reunión cuando la pregunta no va con él: mira al techo. El<bos>es el techo de la sala, y una columna encendida ahí no es un hallazgo: es el estado de reposo. - Bloques y saltos. Lo interesante son las celdas lejos de la diagonal: la jugada 12 mirando a la 5, que es cuando se movió la pieza que ahora se captura. Algunas cabezas de capas intermedias hacen eso, y es lo más cerca que se está de ver el modelo «razonando» sobre el tablero.
- Diferencias entre capas. Lo que se cuenta habitualmente es que las primeras capas tienen patrones locales y regulares, las del medio los más estructurados y las últimas los más difusos, porque ya trabajan sobre representaciones muy mezcladas. Es una expectativa razonable y conviene comprobarla antes de repetirla: en este modelo la mitad de esa frase no se cumple.
La jugada 1 (<bos>) reparte su atención sobre todo en: 1 <bos> (100,0 %)
| Jugada que mira | Jugada 1, <bos> | Jugada 2, <w1800> | Jugada 3, <b1800> | Jugada 4, e2e4 | Jugada 5, e7e5 | Jugada 6, g1f3 | Jugada 7, b8c6 | Jugada 8, f1c4 | Jugada 9, d7d6 | Jugada 10, b1c3 | Jugada 11, c8g4 | Jugada 12, f3e5 | Jugada 13, g4d1 | Jugada 14, c4f7 | Jugada 15, e8e7 | Jugada 16, c3d5 |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
1 <bos> | 100,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
2 <w1800> | 99,9 % | 0,1 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
3 <b1800> | 1,7 % | 87,7 % | 10,6 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
4 e2e4 | 0,2 % | 0,2 % | 98,7 % | 0,8 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
5 e7e5 | 0,0 % | 0,2 % | 0,7 % | 0,9 % | 98,2 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
6 g1f3 | 0,0 % | 0,0 % | 0,1 % | 0,0 % | 99,9 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
7 b8c6 | 0,0 % | 0,1 % | 0,4 % | 0,3 % | 2,3 % | 4,9 % | 91,9 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
8 f1c4 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,1 % | 0,0 % | 99,9 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
9 d7d6 | 0,0 % | 0,0 % | 0,5 % | 0,1 % | 0,5 % | 0,7 % | 73,4 % | 19,7 % | 5,1 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
10 b1c3 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 1,2 % | 0,0 % | 98,8 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
11 c8g4 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,5 % | 0,0 % | 99,4 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
12 f3e5 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 100,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
13 g4d1 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 100,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % |
14 c4f7 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 100,0 % | 0,0 % | 0,0 % | 0,0 % |
15 e8e7 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 100,0 % | 0,0 % | 0,0 % |
16 c3d5 | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 0,0 % | 7,0 % | 0,0 % | 93,0 % | 0,0 % |
0100,0 %Subdiagonal encendida (la celda justo a la izquierda de la diagonal): cada jugada mira a la anterior. Diagonal: cada jugada se mira a sí misma. Columna encendida: una jugada a la que mira todo el mundo.
Lo que hay de verdad en este attention.json. Lo que sigue no es lo que se ve «a ojo» en el
mapa: es el resultado de leer las 96 matrices de 16 × 16 con un script de veinte líneas, porque
mirar colores es precisamente la forma de contarse una película. Salen cuatro cosas: una limpia,
una que contradice el punto de arriba, una que lo matiza y una que hay que dejar sin concluir.
| Qué | Dónde | Medido | Lectura |
|---|---|---|---|
| Cabeza de «jugada anterior» | L9H4 | 0,80 de media en la subdiagonal; 1,000 en cuatro filas seguidas. Siguientes: L8H7 0,65, L7H7 0,60, L4H3 0,58, L9H7 0,53; el resto < 0,45 | Limpia: hay una y es inequívoca |
| Primera capa sin patrones | L0H0-L0H7 | KL frente a la uniforme 0,23-0,34 (media del modelo 0,72; L9H4 1,90); en L0H0 la celda más alta de la última fila es 0,198 | Contradice «las primeras capas son locales»: la capa 0 reparte |
Sumideros, y no solo en <bos> |
capas 10-11; L3H7 | 7 de 8 cabezas de la capa 11 con el máximo en <bos> (0,345-0,513); de media, <bos> 0,229 y <b1800> 0,226; L3H7 0,746-0,986 sobre <b1800> |
Estado de reposo, repartido entre dos tokens de control |
| Saltos lejos de la diagonal | L5H5, L4H0, L1H3 | 0,999 sobre d7d6 en la fila de f3e5; 0,987 sobre f1c4 en la de g4d1; el mismo token como máximo en 3 y en 8 filas |
Sin concluir: «columna enganchada» o «pieza que se captura», y una partida no lo distingue |
Lo limpio: hay una cabeza de jugada anterior, y es inequívoca. Es la capa 9, cabeza 4
(L9H4), con una media de 0,80 de peso sobre la celda inmediatamente a la izquierda de la
diagonal, frente a un máximo de 0,65 en las otras noventa y cinco. Y no es una media de valores
tibios: en 12 de sus 15 filas el máximo cae justo ahí, con 0,988 (b1c3 mirando a d7d6), 0,994 y
luego 1,000, 1,000, 1,000, 1,000 en las filas de f3e5, g4d1, c4f7 y e8e7. Esa cabeza no
hace otra cosa. Detrás van L8H7 (0,65 de media), L7H7 (0,60), L4H3 (0,58) y L9H7 (0,53); las 91
restantes se quedan por debajo de 0,45. Las filas en las que L9H4 se sale del guion son las tres
primeras jugadas negras: dos (e7e5 y b8c6) ponen el peso en sí mismas (0,982 y 0,919 sobre
la diagonal) y d7d6 lo pone en b8c6 (0,734). No sé por qué, y prefiero dejarlo escrito así a
inventarme una explicación.
Lo que contradice el punto de arriba: la capa 0 no tiene patrones locales; no tiene ninguno. Sus
ocho cabezas son casi indistinguibles entre sí y casi uniformes. Midiendo cada fila contra la
distribución uniforme sobre su propio pasado, la divergencia KL de las ocho cabezas de la capa 0 va
de 0,23 a 0,34, cuando la media del modelo entero es 0,72 y L9H4 llega a 1,90; en la última fila, la
celda más alta de L0H0 es el <bos> con 0,198 cuando lo uniforme serían 0,062. La primera capa
reparte, y ya está. La otra mitad de la expectativa sí se cumple: las capas 10 y 11 vuelven a ser
difusas —la atención a la jugada anterior cae a 0,07-0,17— y vuelcan la masa en los tokens de
control; en la capa 11, siete de las ocho cabezas tienen el máximo de su última fila en el
<bos>, entre 0,345 y 0,513, y la octava en <b1800> con 0,570. Con un matiz que el punto del
sumidero no anticipaba: aquí el sumidero no es solo el <bos>. Promediando el modelo entero sobre
las filas de jugada, el <bos> se lleva 0,229 de la atención y <b1800> —el último token de
control, el que va justo antes de la primera jugada— 0,226, prácticamente lo mismo. Y hay
cabezas dedicadas a él en exclusiva: L3H7 pone entre 0,746 y 0,986 sobre <b1800> en ocho filas
seguidas.
Y lo que hay que dejar sin concluir, que era lo que más ganas había de encontrar. Sí aparecen celdas
lejos de la diagonal con pesos altísimos, y son tentadoras: L5H5 pone 0,999 sobre d7d6 justo en
la fila en la que toca escribir f3e5, que es el caballo capturando en e5 después de que ese d7d6
dejara de defender el peón; L4H0 pone 0,987 sobre f1c4 en la fila de g4d1. Es exactamente la
película que el aviso de aquí abajo pide no contarse, y los números explican por qué: esas mismas
cabezas apuntan al mismo token en varias filas seguidas. L5H5 tiene su máximo en d7d6 en tres
filas distintas, y L1H3 lo tiene en e7e5 en ocho de las suyas. Eso se parece mucho más a «esta
cabeza se ha enganchado a una columna» que a «esta cabeza busca la pieza que se captura», y con
una partida de 16 tokens no hay manera de distinguir las dos cosas. Para distinguirlas harían
falta un centenar de partidas y una estadística sobre ellas; la isla es un instrumento de
depuración, no un experimento.
Por qué el modelo «aprende el tablero» sin verlo
Al terminar esta sección sabrás qué se ha demostrado exactamente sobre los modelos de lenguaje de ajedrez y sus representaciones internas, y —más importante para no hacer el ridículo— qué no.
El hecho que sostiene este módulo entero es raro cuando te paras a pensarlo. El modelo nunca ha visto un tablero. No tiene casillas, ni piezas, ni reglas; tiene 2 030 símbolos y seis millones de secuencias. Y acaba proponiendo jugadas legales el 99,4 % de las veces, lo que exige saber dónde está cada pieza después de cuarenta movimientos, incluida la torre que no se mueve desde la jugada 3. Esa información no está en el token actual: está repartida por toda la historia y hay que reconstruirla.
La lectura obligatoria de este módulo es el trabajo de Adam Karvonen, Chess-GPT’s Internal World Model (2024). Es corto, es claro y es la referencia que usa el curso para calibrar expectativas: un nanoGPT de 50 millones de parámetros, entrenado sobre 16 millones de partidas de Lichess en PGN a nivel de carácter, alcanzó alrededor de 1300 de Elo con un 99,8 % de jugadas legales en un día de GPU. Léelo en adamkarvonen.github.io; el código y los modelos están publicados con licencia MIT.
Conviene poner esa cifra al lado de la nuestra sin maquillarla: 1300 de Elo y 99,8 % de legalidad,
frente a los 1359 y 99,4 % que mide small. Karvonen usó 16 millones de partidas y un día de
GPU; nosotros, 5,9 millones y 42 minutos, con un modelo un 22 % más pequeño. El curso no está
igualando ese resultado, y la lección no finge que sí.
Qué se demostró exactamente. El método son sondas lineales (linear probes). Se congela el modelo, se toman las activaciones de una capa intermedia en una posición concreta de la secuencia, y se entrena una regresión lineal —nada más: una matriz, sin capas ocultas— para predecir el estado del tablero, casilla a casilla, a partir de esas activaciones. La sonda acierta qué pieza hay en cada una de las 64 casillas con una exactitud altísima. Es decir: el estado del tablero está presente en las activaciones del modelo de forma linealmente legible, aunque nadie se lo enseñara y aunque su entrada sean caracteres de texto. Se entrenaron también sondas para la fuerza estimada de los jugadores, con resultados análogos.
Una sonda lineal es una regla de tres sobre las activaciones: tantas unidades de esta dimensión, tantas de aquella, sumadas con unos coeficientes fijos, dan «hay un caballo en f3». Ninguna capacidad de cálculo, ningún «razonamiento» añadido por la sonda. Por eso el resultado es fuerte: si una regla de tres ya lee el tablero, es que el tablero no estaba escondido en las activaciones, estaba escrito. Una sonda con capas ocultas podría estar calculándolo ella, y no demostraría nada sobre el modelo.
Y hay una segunda parte, que es la que convierte una correlación en algo más fuerte. Karvonen interviene: toma la dirección que la sonda asocia a «hay una pieza en esta casilla» y modifica las activaciones del modelo a lo largo de esa dirección, como si borrase una pieza del tablero interno. Las jugadas que el modelo propone después cambian de forma coherente con el tablero editado. Eso descarta la explicación aburrida de que la sonda esté leyendo algo incidental: la representación no solo está ahí, sino que el modelo la usa.
Qué no se demostró, y esto importa más para tu criterio:
- No se demostró que el modelo tenga reglas. Que pueda reconstruir el tablero no implica que represente «el alfil se mueve en diagonal»; puede haber aprendido las regularidades estadísticas que el tablero produce, sin la regla que las genera. Las dos cosas son distintas, y la segunda generaliza a posiciones raras mientras que la primera no.
- No se demostró que busque ni calcule. No hay nada parecido a una exploración de variantes: hay una única pasada hacia delante con un número fijo de operaciones. Ese es, seguramente, el techo de este enfoque, y es el motivo de que 2895 de Elo se consiguieran con 270 millones de parámetros y diez millones de partidas anotadas por Stockfish, no con trucos.
- No se demostró que cualquier representación interna sea recuperable. Una sonda lineal encuentra lo que es linealmente legible; que no encuentres algo no significa que no esté, solo que no está de esa forma. El resultado negativo de una sonda no es un resultado.
- No se demostró nada sobre nuestra tokenización. Karvonen usó PGN por carácter; Rukh usa una jugada por token. Es razonable esperar un comportamiento parecido o mejor —al modelo se le ahorra reconstruir la noción de jugada— pero es una hipótesis hasta que la midas, y por eso M3 monta sondas lineales sobre nuestro propio encoder en vez de citar el artículo.
- Y una trampa lógica que conviene nombrar: el estado del tablero es una función determinista de la lista de jugadas, así que la información está necesariamente en la entrada. Lo sorprendente no es que esté, sino que el modelo la haya destilado en una forma explícita y lineal; esa es la frontera exacta entre «memoriza la secuencia» y «modela el dominio».
Qué has aprendido
Cómo sale el modelo de PyTorch y llega al navegador, y las cuatro comprobaciones que hacen que eso sea una afirmación y no una esperanza: el grafo que devuelve solo el último paso, el eje dinámico comprobado ejecutando, los metadatos que viajan dentro del fichero y la paridad medida sobre posiciones reales.
Cómo se mide: uv run rukh export --ckpt checkpoints/small/best.pt --out artifacts/onnx/small --fp16 --int8 --check-parity imprime las tres paridades. Medido: 100 % en fp32, 99,80 % en
fp16 y 95,40 % en int8, con el listón del hito en 99,9 % para los tres. Dos de los tres no lo
pasan, y está escrito así.
Lo siguiente es el navegador: el worker que posee la sesión de ONNX Runtime, el contrato que el fichero tiene que cumplir antes de que se le deje jugar, y la barra de progreso que se mueve de verdad.