rukh · lab

// 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.

  • tiempo de trabajo190 min
  • además, ejecución sin supervisión+ 20 min de GPU y red
  • nivel medio
  • actualizado el22 de septiembre de 2026

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

LabRutaRelojDejaAtajo
Exportar a ONNXLos tres ONNX con su comprobación de paridad, y los JSON de las dos islascuesta máquina~12 min de exportación + minutos de los dos scriptsmodel{,-fp16,-int8}.onnx, parity.json, attention.jsonlos 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:

pyproject.toml
"onnx>=1.23",
"onnxscript>=0.6",
"onnxruntime>=1.30",
"onnxconverter-common>=1.16",
"safetensors>=0.6",

pyproject.tomllíneas 25-29 · p2

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

src/rukh/export/onnx.py
"""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 the
model 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 browser
from slicing a megabyte of logits per move.
``torch.onnx.export`` is tried with the dynamo exporter first (the default since torch 2.9 and
what ``docs/spec/02`` asks for) and falls back to the legacy TorchScript tracer with a warning
when 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 when
it is known which one ran.
``dynamic_seq`` is not taken on trust. The legacy tracer happily bakes the traced length into
the graph while still being asked for a dynamic axis, and the demo feeds a sequence that grows
by one token per move, so the exported file is **run** at two different lengths and the flag
reports what actually worked. The model's context (``block``) travels with the file as ONNX
metadata, so the browser knows the limit without being told separately.
"""

src/rukh/export/onnx.pylíneas 1-19 · p2

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.

src/rukh/export/onnx.py
from __future__ import annotations
import contextlib
import logging
import sys
from collections.abc import Iterator
from pathlib import Path
from typing import Any, Literal
import torch
from pydantic import BaseModel, ConfigDict
from 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 = 18
METADATA_PREFIX = "rukh_"

src/rukh/export/onnx.pylíneas 21-45 · p2

src/rukh/export/onnx.py
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, :]

src/rukh/export/onnx.pylíneas 48-57 · p2

Diez líneas, y son la mitad del rendimiento de la demo.

src/rukh/export/onnx.py
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 = None

src/rukh/export/onnx.pylíneas 60-81 · p2

ExportResult 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.

src/rukh/export/onnx.py
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_NAME

src/rukh/export/onnx.pylíneas 84-87 · p2

src/rukh/export/onnx.py
def _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 axes

src/rukh/export/onnx.pylíneas 90-106 · p2

Dos 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.

src/rukh/export/onnx.py
@contextlib.contextmanager
def _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)

src/rukh/export/onnx.pylíneas 109-130 · p2

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.

src/rukh/export/onnx.py
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)

src/rukh/export/onnx.pylíneas 133-155 · p2

src/rukh/export/onnx.py
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,
)

src/rukh/export/onnx.pylíneas 157-186 · p2

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.

src/rukh/export/onnx.py
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,
)

src/rukh/export/onnx.pylíneas 187-221 · p2

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.

src/rukh/export/onnx.py
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})

src/rukh/export/onnx.pylíneas 224-227 · p2

src/rukh/export/onnx.py
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 True

src/rukh/export/onnx.pylíneas 230-249 · p2

Las 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.

src/rukh/export/onnx.py
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 entries

src/rukh/export/onnx.pylíneas 252-273 · p2

Los 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.

src/rukh/export/onnx.py
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}

src/rukh/export/onnx.pylíneas 276-288 · p2

src/rukh/export/onnx.py
def _load(ckpt: Path) -> MoveDecoder:
from rukh.train import load_model
model, _payload = load_model(ckpt)
return model

src/rukh/export/onnx.pylíneas 291-295 · p2

quantize.py: la mitad y la cuarta parte

src/rukh/export/quantize.py
"""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 int8
one (about 40 MB), which is the difference between a demo that loads on a phone and one that
does not. Quantization is restricted to ``MatMul`` and ``Gemm``: those are the weights that make
up 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 and
outputs in float32 and leaves the ops that overflow in fp32). When it is not installed there is
a 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.
"""

src/rukh/export/quantize.pylíneas 1-12 · p2

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.

src/rukh/export/quantize.py
from __future__ import annotations
import logging
from pathlib import Path
from 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"]

src/rukh/export/quantize.pylíneas 14-26 · p2

src/rukh/export/quantize.py
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.0

src/rukh/export/quantize.pylíneas 29-42 · p2

src/rukh/export/quantize.py
def _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 / name

src/rukh/export/quantize.pylíneas 45-49 · p2

src/rukh/export/quantize.py
def 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,
)

src/rukh/export/quantize.pylíneas 52-76 · p2

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.

src/rukh/export/quantize.py
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 model

src/rukh/export/quantize.pylíneas 79-115 · p2

El 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.

src/rukh/export/quantize.py
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,
)

src/rukh/export/quantize.pylíneas 118-137 · p2

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?

src/rukh/export/parity.py
"""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 on
0.1 % of positions is fine and one that agrees on the numbers but not on the move is not. Both
are 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 is
what tells fp32, fp16 and int8 apart.
The positions are the thousand **validation** prefixes of ``docs/spec/02`` §6, the same ones the
legality and accuracy metrics use: a random legal walk visits positions no human would reach, so
agreeing on them says little about the file the demo will load. ``random_prefixes`` stays for the
case where no validation parquet is around (and for the tests), and which of the two was used is
reported alongside the number.
"""

src/rukh/export/parity.pylíneas 1-15 · p2

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.

src/rukh/export/parity.py
from __future__ import annotations
import logging
import random
from collections.abc import Sequence
from pathlib import Path
import chess
import numpy as np
import torch
from pydantic import BaseModel, ConfigDict
from rukh.export.onnx import INPUT_NAME, LastStepLogits
from rukh.models import MoveDecoder
from rukh.tokenize.uci_vocab import UciTokenizer
log = logging.getLogger(__name__)
DEFAULT_N = 1_000
MAX_MISMATCHES = 20
DEFAULT_GAMES = "data/uci/year=2025/month=02/games.parquet"
"""The validation month of P1: the parity positions come from here when it exists."""

src/rukh/export/parity.pylíneas 17-38 · p2

src/rukh/export/parity.py
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."""

src/rukh/export/parity.pylíneas 41-50 · p2

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.

src/rukh/export/parity.py
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 prefixes

src/rukh/export/parity.pylíneas 53-83 · p2

src/rukh/export/parity.py
def 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]

src/rukh/export/parity.pylíneas 86-97 · p2

src/rukh/export/parity.py
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", warning

src/rukh/export/parity.pylíneas 100-120 · p2

Las 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.

src/rukh/export/parity.py
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"])

src/rukh/export/parity.pylíneas 123-127 · p2

src/rukh/export/parity.py
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,
)

src/rukh/export/parity.pylíneas 130-162 · p2

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.

src/rukh/export/parity.py
def _load(ckpt: Path) -> MoveDecoder:
from rukh.train import load_model
model, _payload = load_model(ckpt)
return model

src/rukh/export/parity.pylíneas 165-169 · p2

__init__.py: rukh export en una función

src/rukh/export/__init__.py
"""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 int8
files the two browser backends need, and check every file it produced against PyTorch.
"""

src/rukh/export/__init__.pylíneas 1-5 · p2

src/rukh/export/__init__.py
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,
)

src/rukh/export/__init__.pylíneas 7-43 · p2

src/rukh/export/__init__.py
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 = None

src/rukh/export/__init__.pylíneas 46-58 · p2

src/rukh/export/__init__.py
def 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 bundle

src/rukh/export/__init__.pylíneas 61-106 · p2

export_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_props de 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.
src/rukh/export/__init__.py
__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",
]

src/rukh/export/__init__.pylíneas 109-137 · p2

El comando, y lo que midió

src/rukh/cli.py
@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}")

src/rukh/cli.pylíneas 499-576 · p2

Terminal
uv run rukh export --ckpt checkpoints/small/best.pt --out artifacts/onnx/small --fp16 --int8 --check-parity

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

exporter: dynamo (opset 18)
kind: decoder -> logits
shapes: batch dynamic=True, sequence dynamic=True (verified), block=200
fp32: 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
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:

  1. 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.
  2. 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.
  3. 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, exactitud top-1 y legalidad a lo largo del entrenamiento20 checkpoints, del paso 1000 al 20000. Debajo de las curvas, una barra por checkpoint con su legalidad sin máscara frente al listón del 99 %. Los valores exactos del checkpoint seleccionado están en la lista que sigue al gráfico.2,751,361,502,002,5052,6 %32,0 %35 %40 %45 %50 %paso 1 000paso 20 0005 00010 00015 000legalidad sin máscara (argmax)100 %85 %paso 1 000: 88,0 %paso 2 000: 94,0 %paso 3 000: 95,0 %paso 4 000: 97,8 %paso 5 000: 97,0 %paso 6 000: 97,0 %paso 7 000: 98,5 %paso 8 000: 99,0 %paso 9 000: 98,8 %paso 10 000: 99,3 %paso 11 000: 99,5 %paso 12 000: 99,3 %paso 13 000: 99,0 %paso 14 000: 99,3 %paso 15 000: 99,8 %paso 16 000: 99,3 %paso 17 000: 99,5 %paso 18 000: 99,3 %paso 19 000: 99,3 %paso 20 000: 99,3 %listón 99 %
Eje izquierdo: pérdida (entrenamiento y validación). Eje derecho: top-1 de siguiente jugada en validación. Tira inferior: legalidad sin máscara por argmax de cada checkpoint, con el listón del 99 % marcado. La línea vertical marca el checkpoint seleccionado.

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

small-20260919-062911 · preset small · 20 checkpoints · generado 2026-09-19T08:49:05+00:00

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 %)

Capa 9, cabeza 4. Cada fila es una jugada que mira; cada columna, una jugada mirada. El color es relativo al máximo de esta cabeza (100,0 %) y el valor exacto está en el título de cada celda.
Jugada que miraJugada 1, <bos>Jugada 2, <w1800>Jugada 3, <b1800>Jugada 4, e2e4Jugada 5, e7e5Jugada 6, g1f3Jugada 7, b8c6Jugada 8, f1c4Jugada 9, d7d6Jugada 10, b1c3Jugada 11, c8g4Jugada 12, f3e5Jugada 13, g4d1Jugada 14, c4f7Jugada 15, e8e7Jugada 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 e2e40,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 e7e50,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 g1f30,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 b8c60,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 f1c40,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 d7d60,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 b1c30,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 c8g40,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 f3e50,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 g4d10,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 c4f70,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 e8e70,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 c3d50,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.

rukh-small · checkpoint checkpoints\small-20260919-062911\best.pt · generado 2026-09-19T08:27:28+00:00

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.