rukh · lab

// M3 · lección 04

Cabezas y casillas: tres preguntas y una segunda entrada

`models/heads.py` y `models/squares.py` enteros: por qué las tres cabezas son una capa lineal, la pérdida conjunta con pesos y máscara, el FEN convertido en 69 tokens fijos, el vocabulario con hash y el test de ida y vuelta que caza una transposición del tablero.

  • tiempo de trabajo100 min
  • nivel medio
  • actualizado el23 de septiembre de 2026

Lección 4 de 11 del módulo «El encoder». Viene de «El encoder a mano» y sigue en «Masked move modeling».

Qué vas a construir

Dos ficheros que cierran el modelo: las tres cabezas que leen el vector agrupado (models/heads.py) y el traductor que convierte un FENFENCadena de texto que describe una posición completa: piezas por fila, turno, derechos de enroque, casilla al paso y contadores de jugadas. Es la clave con la que se cruzan las posiciones de las partidas con las evaluaciones públicas de Stockfish. en los 69 tokens de la segunda representación de entrada (models/squares.py). Con esto solo falta entrenarlo.

Las tres cabezas

src/rukh/models/heads.py
"""The three supervised heads of M3, all of them reading one vector: ``PositionEncoder.pool``.
The heads are deliberately tiny — one linear layer each — because that is what makes the
experiment mean something. If a linear probe on the pooled representation can say how good a
position is, whether the last move threw the game away and who is going to win, then the
representation already contains those facts and masked move modeling put them there. A deep
head would be able to learn them by itself and would tell us nothing about the encoder.
* ``ValueHead``: one scalar through ``tanh``, matching the ``tanh(cp / 400)`` label, so the
output is bounded and a mate cannot dominate the loss.
* ``BlunderHead``: one logit, because the label is missing for a good part of the rows (a
position whose predecessor is not in the table) and a masked binary cross entropy is the
honest way to skip them.
* ``ResultHead``: three classes, White / draw / Black.
``MultiHead`` trains them together over a shared encoder, with a weight per head: the three
tasks have different scales (an MSE in ``[0, 4]``, two cross entropies) and different amounts of
data, so "just add them up" would silently make one of them the only one that matters.
"""
from __future__ import annotations
from typing import Literal
import torch
from torch import Tensor, nn
from torch.nn import functional as F
from rukh.config import BaseConfig
from rukh.models.encoder import PositionEncoder
HEADS = ("value", "blunder", "result")
RESULT_CLASSES = 3

src/rukh/models/heads.pylíneas 1-33 · p3

El docstring es el diseño del experimento: las cabezas son pequeñas a propósito, la sonda más tonta que funciona de la lección 1. HEADS es una tupla que fija el orden de las cabezas al sumar la pérdida y al imprimir métricas, y es el único sitio donde está escrito cuántas hay.

src/rukh/models/heads.py
class HeadWeights(BaseConfig):
"""Weight of each head in the joint loss; ``0`` switches a head off without removing it."""
value: float = 1.0
blunder: float = 1.0
result: float = 0.5

src/rukh/models/heads.pylíneas 36-41 · p3

Hereda de BaseConfig, así que una clave mal escrita en el YAML —results: 0.5— es un error al cargar y no un peso que nunca se aplicó. La cabeza de resultado pesa la mitad porque es la más ruidosa: su etiqueta es de la partida entera, repetida en todas sus posiciones.

Hay pesos porque las tres pérdidas viven en escalas distintas —un error cuadrático pequeño frente a dos entropías cruzadas— y tienen cantidades de datos distintas. Sumarlas a pelo deja que la de números más grandes decida sola hacia dónde se mueve el tronco. Un peso explícito, aunque valga 1, es una decisión visible, y el 0 permite medir qué aporta cada tarea sin cambiar la forma del checkpoint. Es el problema de cualquier pérdida con varios términos, como la recompensa y la penalización de KL de M5.

src/rukh/models/heads.py
class ValueHead(nn.Module):
"""How good is this position for White: a scalar in ``(-1, 1)``."""
def __init__(self, d_model: int) -> None:
super().__init__()
self.proj = nn.Linear(d_model, 1)
def forward(self, pooled: Tensor) -> Tensor:
return torch.tanh(self.proj(pooled)).squeeze(-1)
class BlunderHead(nn.Module):
"""Did the move that led here throw the game away: one logit (not a probability)."""
def __init__(self, d_model: int) -> None:
super().__init__()
self.proj = nn.Linear(d_model, 1)
def forward(self, pooled: Tensor) -> Tensor:
return self.proj(pooled).squeeze(-1)
class ResultHead(nn.Module):
"""How the game ended: logits over White / draw / Black."""
def __init__(self, d_model: int) -> None:
super().__init__()
self.proj = nn.Linear(d_model, RESULT_CLASSES)
def forward(self, pooled: Tensor) -> Tensor:
return self.proj(pooled)

src/rukh/models/heads.pylíneas 44-74 · p3

Dos detalles importan. ValueHead pasa por tanh para vivir en el mismo rango que su etiqueta, tanh(cp / 400). BlunderHead emite un logit sin acotar, porque binary_cross_entropy_with_logits es numéricamente estable donde una sigmoide seguida de un logaritmo no lo es: con una probabilidad de 1e-8, el logaritmo se va al infinito en float32. La demo aplica el torch.sigmoid fuera.

El otro es el .squeeze(-1). nn.Linear(d, 1) devuelve (B, 1) y la etiqueta es (B,). Sin el squeeze, F.mse_loss no falla: difunde los dos tensores a (B, B) y promedia todas las parejas, incluidas las que emparejan la predicción de una posición con la etiqueta de otra. La pérdida sale finita, el entrenamiento arranca y lo que se minimiza es otra cosa.

src/rukh/models/heads.py
class MultiHead(nn.Module):
"""An encoder plus the three heads, trained on one pooled representation."""
def __init__(
self,
encoder: PositionEncoder,
weights: HeadWeights | None = None,
pooling: Literal["cls", "mean"] = "mean",
) -> None:
super().__init__()
self.encoder = encoder
self.weights = weights or HeadWeights()
self.pooling = pooling
d_model = encoder.cfg.d_model
self.value = ValueHead(d_model)
self.blunder = BlunderHead(d_model)
self.result = ResultHead(d_model)
def pooled(self, idx: Tensor, attention_mask: Tensor | None = None) -> Tensor:
"""The position's representation, ``(B, d_model)``."""
return self.encoder.pool(idx, self.pooling, attention_mask)
def forward(self, idx: Tensor, attention_mask: Tensor | None = None) -> dict[str, Tensor]:
"""``{"value": (B,), "blunder": (B,) logits, "result": (B, 3) logits}``."""
return self.from_pooled(self.pooled(idx, attention_mask))
def from_pooled(self, pooled: Tensor) -> dict[str, Tensor]:
"""The three heads' outputs for an already pooled batch."""
return {
"value": self.value(pooled),
"blunder": self.blunder(pooled),
"result": self.result(pooled),
}

src/rukh/models/heads.pylíneas 77-109 · p3

MultiHead contiene al encoder, así que el tronco aparece con el prefijo encoder. en los checkpoints de las cabezas. Cargar un encoder preentrenado en un MultiHead es añadir ese prefijo, y olvidarlo es entrenar cabezas sobre pesos aleatorios creyendo que están sobre un preentrenamiento.

La separación entre forward y from_pooled hace que el vector se calcule una vez para las tres cabezas, y es lo que usa el exportador de la lección 9 para meter el pooling una sola vez en el grafo ONNX. forward devuelve un diccionario: con tres salidas de formas distintas, una tupla convierte cada llamada en un ejercicio de memoria sobre el orden.

src/rukh/models/heads.py
def loss(
self, outputs: dict[str, Tensor], targets: dict[str, Tensor]
) -> tuple[Tensor, dict[str, Tensor]]:
"""``(total, per head)``; a head with no labelled row in the batch contributes zero.
``targets`` carries ``value`` (float), ``blunder`` (0/1 float) with its ``blunder_mask``
(the rows that have a label at all) and ``result`` (class index).
"""
zero = torch.zeros((), device=outputs["value"].device, dtype=outputs["value"].dtype)
parts = {
"value": F.mse_loss(outputs["value"], targets["value"].to(outputs["value"].dtype)),
"blunder": zero,
"result": F.cross_entropy(outputs["result"], targets["result"].long()),
}
mask = targets["blunder_mask"].bool()
if bool(mask.any()):
parts["blunder"] = F.binary_cross_entropy_with_logits(
outputs["blunder"][mask], targets["blunder"][mask].to(outputs["blunder"].dtype)
)
total = sum(
(getattr(self.weights, name) * parts[name] for name in HEADS),
start=torch.zeros_like(zero),
)
return total, parts

src/rukh/models/heads.pylíneas 111-134 · p3

Tres decisiones:

  • Devuelve el total y las partes. Con un tronco compartido, una tarea puede mejorar a costa de otra sin que el agregado lo diga. La lección 11 lo enseña: al descongelar el modelo entero, la pérdida de la cabeza de errores empeora, y con solo el total sería invisible.
  • La máscara de la cabeza de errores. outputs["blunder"][mask] usa solo las filas con etiqueta; las demás siguen valiendo para el valor y el resultado. Con una clase positiva del 3,72 %, etiquetar «no lo sé» como «no fue error» envenenaría la clase mayoritaria con las filas de las que hay menos evidencia.
  • El if bool(mask.any()) y el zero. Un lote sin filas etiquetadas daría una entropía cruzada sobre tensores vacíos, que es nan, y ese nan llegaría por la suma a los gradientes de las tres cabezas y del tronco. Un solo lote desafortunado arruinaría la tirada; con el zero, esa cabeza aporta cero ese paso.

El .to(outputs["value"].dtype) de las etiquetas es la concesión a bf16bfloat16Formato de coma flotante de 16 bits con los mismos 8 bits de exponente que fp32 y solo 7 de mantisa: pierde precisión pero conserva el rango, así que un gradiente pequeño no se va a cero. Por eso entrena sin escalado de pérdida, al contrario que fp16, cuyo exponente de 5 bits desborda por abajo.: bajo autocast la salida es bfloat16 y F.mse_loss con tipos distintos lanza. Se convierte la etiqueta y no la salida, que sacaría el cálculo del autocast.

src/rukh/models/heads.py
def head_parameters(self) -> list[nn.Parameter]:
"""Every parameter of the three heads, the ones a ``probe`` run is allowed to move."""
return [
param for head in (self.value, self.blunder, self.result) for param in head.parameters()
]

src/rukh/models/heads.pylíneas 136-140 · p3

Existe para que la lección 6 pueda escribir «descongela esto» en una línea.

Los tests de las cabezas

tests/unit/test_heads.py
"""Tests for the three heads and their staged fine-tuning: gradients, freezing and the curve."""
from __future__ import annotations
from pathlib import Path
import polars as pl
import pytest
import torch
from rukh.models import EncoderConfig, MultiHead, PositionEncoder
from rukh.models.heads import HEADS, BlunderHead, HeadWeights, ResultHead, ValueHead
from rukh.models.squares import SQUARE_TOKENS, fen_to_tokens
from rukh.train import HeadsConfig, freeze_encoder, label_curve, train_heads
from rukh.train.checkpoint import load_checkpoint, save_checkpoint
from rukh.train.heads import LabelledPositions, build_frames, set_training_mode
pytestmark = pytest.mark.unit
TOY = EncoderConfig(input="squares", n_layer=3, n_head=2, d_model=32, dropout=0.0)
# Four boards whose evaluation the encoder can actually read off the pieces: the full set, the
# same without Black's queen, without White's queen, and without either. A toy task has to be
# learnable or "the loss went down" measures nothing.
BOARDS = (
"rnbqkbnr/pppppppp/8/8/8/8/PPPPPPPP/RNBQKBNR",
"rnb1kbnr/pppppppp/8/8/8/8/PPPPPPPP/RNBQKBNR",
"rnbqkbnr/pppppppp/8/8/8/8/PPPPPPPP/RNB1KBNR",
"rnb1kbnr/pppppppp/8/8/8/8/PPPPPPPP/RNB1KBNR",
)
BOARD_CP = (0, 300, -300, 0)
BOARD_RESULT = ("1/2-1/2", "1-0", "0-1", "1/2-1/2")
WHITE_TO_MOVE = f"{BOARDS[0]} w KQkq -"

tests/unit/test_heads.pylíneas 1-32 · p3

El comentario de las cuatro posiciones vale para cualquier test de entrenamiento: una tarea de juguete tiene que ser aprendible o «la pérdida bajó» no mide nada. Un modelo que mire las piezas puede aprender que sin la dama negra se evalúa +300; con etiquetas aleatorias, la pérdida también bajaría por memorización y el test pasaría sin comprobar nada.

tests/unit/test_heads.py
def toy_batch(rows: int = 4) -> dict[str, torch.Tensor]:
idx = torch.tensor([fen_to_tokens(WHITE_TO_MOVE) for _ in range(rows)])
return {
"idx": idx,
"value": torch.linspace(-0.9, 0.9, rows),
"blunder": torch.tensor([1.0, 0.0] * (rows // 2)),
"blunder_mask": torch.tensor([True, False] * (rows // 2)),
"result": torch.tensor([0, 1, 2, 0][:rows]),
}

tests/unit/test_heads.pylíneas 45-53 · p3

La mitad de las filas del lote no tienen etiqueta de error, para comprobar la máscara sin montar un dataset.

tests/unit/test_heads.py
def test_each_head_has_the_shape_its_label_needs() -> None:
pooled = torch.randn(5, 32)
value = ValueHead(32)(pooled)
assert value.shape == (5,) and bool((value.abs() < 1.0).all())
assert BlunderHead(32)(pooled).shape == (5,) # a logit, not a probability
assert ResultHead(32)(pooled).shape == (5, 3)
def test_multihead_answers_the_three_questions_at_once() -> None:
model = MultiHead(toy_encoder())
batch = toy_batch()
outputs = model(batch["idx"])
assert set(outputs) == set(HEADS)
assert outputs["value"].shape == (4,)
assert outputs["blunder"].shape == (4,)
assert outputs["result"].shape == (4, 3)
assert model.pooled(batch["idx"]).shape == (4, 32)
assert model.encoder.cfg.seq == SQUARE_TOKENS

tests/unit/test_heads.pylíneas 106-123 · p3

El primero caza la difusión silenciosa del squeeze y, de paso, que el tanh sigue puesto.

tests/unit/test_heads.py
def test_multihead_propagates_the_three_gradients() -> None:
model = MultiHead(toy_encoder(), HeadWeights(value=1.0, blunder=1.0, result=1.0))
batch = toy_batch()
total, parts = model.loss(model(batch["idx"]), batch)
assert set(parts) == set(HEADS)
assert all(torch.isfinite(part) for part in parts.values())
total.backward()
for name in HEADS:
grad = getattr(model, name).proj.weight.grad
assert grad is not None and float(grad.abs().sum()) > 0.0, name
body = model.encoder.blocks[0].mlp.fc.weight.grad
assert body is not None and float(body.abs().sum()) > 0.0 # the encoder gets them too

tests/unit/test_heads.pylíneas 126-137 · p3

Comprueba que la pérdida está conectada. Un detach() de más —un .data, un torch.no_grad() que envuelve demasiado— produce un modelo que entrena sin aprender, y la única señal sería que la pérdida no baja. Aquí el fallo es inmediato y dice qué cabeza. La última aserción exige que el gradiente llegue al tronco, que es lo que distingue un afinado de un probe.

tests/unit/test_heads.py
def test_a_head_with_weight_zero_stops_contributing() -> None:
model = MultiHead(toy_encoder(), HeadWeights(value=1.0, blunder=0.0, result=0.0))
batch = toy_batch()
total, parts = model.loss(model(batch["idx"]), batch)
assert torch.allclose(total, parts["value"], atol=1e-6)
total.backward()
assert float(model.blunder.proj.weight.grad.abs().sum()) == 0.0
assert float(model.value.proj.weight.grad.abs().sum()) > 0.0
def test_a_batch_without_a_blunder_label_costs_nothing_instead_of_nan() -> None:
model = MultiHead(toy_encoder())
batch = toy_batch()
batch["blunder_mask"] = torch.zeros_like(batch["blunder_mask"])
total, parts = model.loss(model(batch["idx"]), batch)
assert float(parts["blunder"]) == 0.0
assert torch.isfinite(total)

tests/unit/test_heads.pylíneas 140-156 · p3

El segundo es el contrato del if bool(mask.any()) escrito como aserción: cero, no nan. Impide que alguien «simplifique» esa rama en un refactor.

squares.py: un FEN en 69 tokens

La otra mitad de la lección es la segunda representación de entrada. El fichero abre con la especificación del formato.

src/rukh/models/squares.py
"""The ``squares`` input scheme: a FEN as 69 fixed tokens, the didactic rival of ``moves``.
The ``moves`` scheme feeds the encoder the game so far (the decoder's own UCI vocabulary); this
one feeds it the position itself, so the two can be compared on the same heads. That comparison
is the point of the module: it is the "what is a good representation?" lesson of M3.
Layout (always exactly ``SQUARE_TOKENS`` = 69 positions, so no padding is ever needed)::
0 <cls> pooling anchor, the counterpart of the decoder's <bos>
1..64 the 64 squares piece or <empty>, file-major (a1, a2, ..., a8, b1, ..., h8)
65 side to move turn:w / turn:b
66 castling rights castle:<KQkq subset>, 16 combinations in one token
67 en passant ep:none or the file, ep:a ... ep:h
68 halfmove clock clock:0 ... clock:3, the bucketed 50-move counter
The square order is ``rukh.tokenize.uci_vocab.squares()``, file-major, the same enumeration the
UCI vocabulary uses for move endpoints: one ordering for the whole project.
The plan's arithmetic (64 + turn + 4 castling + en passant + clock) adds up to 71, not to the 69
it fixes; 69 is the binding number, so the four castling rights travel in one token with 16
values (they are four bits of one fact, and this keeps every other component in a slot of its
own) and the position that frees up became ``<cls>``, which gives ``pool("cls")`` a real anchor
in this scheme instead of borrowing square a1.
Halfmove-clock buckets, chosen so that the boundaries mean something over the board:
=========== ===================================================================
``clock:0`` 0-5 half-moves: a pawn moved or a piece was captured very recently
``clock:1`` 6-24: a normal manoeuvring stretch
``clock:2`` 25-49: the 50-move rule is in sight and shapes the plan
``clock:3`` 50 or more: a draw can be claimed
=========== ===================================================================
A four-field FEN (``fen4``, the key of the P1 positions and evaluations) carries no counters, so
it reads back as ``clock:0``; that is the value every label built from P1 data will have.
``SQUARE_VOCAB`` is fixed by construction, never learned from data, and ``vocab_hash`` pins it:
a change to the enumeration invalidates every encoder trained with it, exactly like the UCI
vocabulary of P1.
"""

src/rukh/models/squares.pylíneas 1-40 · p3

De ese docstring, cuatro cosas son decisiones:

  • El orden por columnas (a1, a2, …, a8, b1, …) es el mismo que usa el vocabulario UCI de M1, así que un índice de casilla significa lo mismo en todo el proyecto. Un FEN se escribe por filas y de la 8 a la 1; la transposición entre las dos convenciones es la parte peligrosa del fichero.
  • Los cuatro derechos de enroque van en un token de 16 valores. La aritmética obvia da 71 posiciones y el formato fija 69; juntar los cuatro bits de un mismo hecho es la elección menos arbitraria, y deja sitio para un <cls> de verdad.
  • Los tramos del reloj tienen fronteras con sentido ajedrecístico («acaba de haber captura», «la regla de las 50 jugadas ya condiciona el plan») en vez de cuartiles, que el modelo no podría interpretar.
  • Las filas de P1 son FEN de cuatro campos, sin contadores, así que el token 68 vale clock:0 en toda la tabla supervisada y no aporta ni un bit. Se queda porque una posición en vivo de la demo sí trae los contadores, y una representación que cambia de forma entre entrenamiento y servicio sería peor que una ranura constante.
src/rukh/models/squares.py
from __future__ import annotations
import hashlib
import json
from rukh.tokenize.uci_vocab import FILES, squares
PIECES = "PNBRQKpnbrqk"
CASTLING_ORDER = "KQkq"
N_SQUARES = 64
SQUARE_TOKENS = 69
"""Length of every ``squares`` sequence: 1 + 64 + 4."""
CLOCK_EDGES = (6, 25, 50)
"""Upper edges (exclusive) of the halfmove-clock buckets; see the module docstring."""
def castling_strings() -> list[str]:
"""The 16 castling combinations in a fixed order: ``-``, ``K``, ``Q``, ``KQ``, ``k``, ..."""
out: list[str] = []
for mask in range(16):
rights = "".join(right for bit, right in enumerate(CASTLING_ORDER) if mask >> bit & 1)
out.append(rights or "-")
return out
def build_square_vocab() -> list[str]:
"""The token list in id order; see the module docstring for the layout it serves."""
tokens = ["<pad>", "<mask>", "<cls>", "<empty>"]
tokens.extend(PIECES)
tokens.extend(["turn:w", "turn:b"])
tokens.extend(f"castle:{rights}" for rights in castling_strings())
tokens.append("ep:none")
tokens.extend(f"ep:{file}" for file in FILES)
tokens.extend(f"clock:{index}" for index in range(len(CLOCK_EDGES) + 1))
return tokens
SQUARE_VOCAB: list[str] = build_square_vocab()
SQUARE_IDS: dict[str, int] = {token: index for index, token in enumerate(SQUARE_VOCAB)}
SQUARE_VOCAB_SIZE = len(SQUARE_VOCAB)
PAD_ID = SQUARE_IDS["<pad>"]
MASK_ID = SQUARE_IDS["<mask>"]
CLS_ID = SQUARE_IDS["<cls>"]
EMPTY_ID = SQUARE_IDS["<empty>"]
CONTROL_IDS: frozenset[int] = frozenset({PAD_ID, MASK_ID, CLS_ID})
"""Never masked by ``apply_masking``: they are the frame of the sequence, not a prediction."""

src/rukh/models/squares.pylíneas 42-89 · p3

El vocabularioVocabularioLa lista de todos los tokens que el modelo conoce, cada uno con un id entero fijo. En Rukh el vocabulario UCI tiene 2 030 entradas: 8 tokens especiales, 54 tramos de Elo y las 1 968 jugadas posibles, enumeradas sin mirar datos. Su tamaño fija el de la capa de embeddings y el de la capa de salida. se construye, no se escribe: una lista de 47 cadenas escrita a mano es una lista que alguien reordena sin querer al añadir un token. <pad> va primero porque es el 0 en los dos esquemas, lo que permite que padding_mask sea idx != 0 sin preguntar por el esquema. CONTROL_IDS es el contrato con el enmascarado de la lección siguiente: los tokens que son el marco de la secuencia y nunca una predicción.

src/rukh/models/squares.py
def vocab_hash() -> str:
"""SHA-256 of the vocabulary as a JSON list: the identity of the ``squares`` scheme."""
payload = json.dumps(SQUARE_VOCAB, ensure_ascii=True, separators=(",", ":"))
return hashlib.sha256(payload.encode("ascii")).hexdigest()
def clock_bucket(halfmove: int) -> int:
"""Index of the halfmove-clock bucket of ``halfmove`` (see ``CLOCK_EDGES``)."""
if halfmove < 0:
raise ValueError(f"the halfmove clock cannot be negative, got {halfmove}")
return sum(halfmove >= edge for edge in CLOCK_EDGES)

src/rukh/models/squares.pylíneas 92-102 · p3

El hash existe por el mismo motivo que el del vocabulario UCI de M1: un encoder solo significa algo con el vocabulario que vio. Si alguien inserta un token en medio de la lista, los ids posteriores se desplazan, los pesos del embedding pasan a describir otras cosas y el modelo sigue cargando y prediciendo. El hash convierte ese cambio en un test rojo; ensure_ascii y separators fijos lo hacen reproducible.

clock_bucket cuenta cuántas fronteras se han pasado: un 30 supera dos y cae en el tramo 2. Sigue siendo correcto si mañana hay cinco fronteras.

src/rukh/models/squares.py
def _placement_ids(placement: str) -> list[int]:
"""The 64 square ids, file-major, from the piece-placement field of a FEN."""
ranks = placement.split("/")
if len(ranks) != 8:
raise ValueError(f"a FEN placement needs 8 ranks, got {len(ranks)}: {placement!r}")
ids = [EMPTY_ID] * N_SQUARES
for row, rank_text in enumerate(ranks):
rank = 7 - row # the placement is written from rank 8 down to rank 1
file = 0
for char in rank_text:
if char.isdigit():
file += int(char)
elif char in PIECES:
if file >= 8:
raise ValueError(f"rank {rank + 1} of {placement!r} is too long")
ids[file * 8 + rank] = SQUARE_IDS[char]
file += 1
else:
raise ValueError(f"unknown piece {char!r} in {placement!r}")
if file != 8:
raise ValueError(f"rank {rank + 1} of {placement!r} covers {file} files, not 8")
return ids

src/rukh/models/squares.pylíneas 105-126 · p3

La función peligrosa del fichero, y el peligro cabe en una expresión: ids[file * 8 + rank]. Escribir rank * 8 + file produce un tablero válido y distinto, la posición reflejada en la diagonal. Ninguna comprobación de forma lo detecta, y el modelo aprende ajedrez en un tablero girado. Es el mismo error que trasponer alto y ancho al cargar imágenes o intercambiar dos columnas de un CSV sin cabecera: los datos siguen teniendo buena pinta.

Las tres validaciones cazan una fila con piezas de más, una fila que no cubre las ocho columnas (una 7 donde debía haber una 8, la errata típica de un FEN escrito a mano) y un carácter que no es ni dígito ni pieza, y todas nombran la fila en el mensaje.

src/rukh/models/squares.py
def _castling_id(field: str) -> int:
"""Token id of a FEN castling field, normalized to the ``KQkq`` order."""
if field in ("-", ""):
return SQUARE_IDS["castle:-"]
unknown = [char for char in field if char not in CASTLING_ORDER]
if unknown:
raise ValueError(f"unsupported castling field {field!r} (Chess960 is out of scope)")
rights = "".join(right for right in CASTLING_ORDER if right in field)
return SQUARE_IDS[f"castle:{rights}"]

src/rukh/models/squares.pylíneas 129-137 · p3

La normalización al orden KQkq hace que qK y Kq sean el mismo token; sin ella, la misma posición escrita por dos generadores de FEN daría dos entradas distintas. El mensaje de error nombra Chess960 porque es el caso que se rechaza: allí el enroque se escribe con la columna de la torre (AHah).

src/rukh/models/squares.py
def fen_to_tokens(fen: str) -> list[int]:
"""The ``SQUARE_TOKENS`` ids of a FEN; four fields (``fen4``) or the full six are accepted.
Raises ``ValueError`` on a malformed FEN: the encoder must never be fed a position that was
silently repaired into a different one.
Token 68, the halfmove clock, is **constant** over every dataset this project builds. A
``fen4`` has no counters, so it reads back as ``clock:0``, and P1's positions table is all
``fen4``: the bucket carries exactly zero information there and the encoder learns a bias
term for it. It stays in the layout because a live position from the demo does have the
counters, and a representation that changes shape between training and serving would be
worse than one slot of constant.
"""
fields = fen.split()
if len(fields) < 4:
raise ValueError(f"a FEN needs at least 4 fields, got {len(fields)}: {fen!r}")
placement, turn, castling, ep = fields[:4]
if turn not in ("w", "b"):
raise ValueError(f"the side to move must be 'w' or 'b', got {turn!r}")
halfmove = int(fields[4]) if len(fields) > 4 else 0
if ep in ("-", ""):
ep_token = "ep:none"
elif len(ep) == 2 and ep[0] in FILES and ep[1] in "36":
ep_token = f"ep:{ep[0]}"
else:
raise ValueError(f"unknown en-passant square {ep!r}")
return [
CLS_ID,
*_placement_ids(placement),
SQUARE_IDS[f"turn:{turn}"],
_castling_id(castling),
SQUARE_IDS[ep_token],
SQUARE_IDS[f"clock:{clock_bucket(halfmove)}"],
]

src/rukh/models/squares.pylíneas 140-173 · p3

La frase que manda está en el docstring: un FEN mal formado lanza. La tentación en un pipeline de datos es reparar y seguir, pero un tablero equivocado que entrena es peor que una ejecución que se para. Se ignora el número de jugada (fields[5]), que se deduce del plyPly (media jugada)Una jugada de un solo bando. 1. e4 e5 son dos plies y una jugada completa. Los filtros del recorte y las longitudes de secuencia del modelo se cuentan en plies porque es lo que ve el modelo: un token por ply.. Y ep[1] in "36" es estricto a propósito: la casilla al paso solo puede estar en la fila 3 o en la 6, y un e4 es un FEN corrupto.

src/rukh/models/squares.py
def tokens_to_strings(tokens: list[int]) -> list[str]:
"""The token strings of a sequence, for debugging and for the lesson's figures."""
return [SQUARE_VOCAB[token] for token in tokens]
def tokens_to_fen(tokens: list[int]) -> str:
"""The four-field FEN a token sequence encodes: the inverse of ``fen_to_tokens``.
The halfmove-clock bucket is dropped because it is lossy by construction, so the round trip
is exact for a ``fen4`` and exact up to the counters for a full FEN.
"""
if len(tokens) != SQUARE_TOKENS:
raise ValueError(f"a squares sequence has {SQUARE_TOKENS} tokens, got {len(tokens)}")
names = tokens_to_strings(tokens)
if names[0] != "<cls>":
raise ValueError(f"a squares sequence starts with <cls>, got {names[0]!r}")
board = names[1 : 1 + N_SQUARES]
ranks: list[str] = []
for rank in range(7, -1, -1):
row, empty = "", 0
for file in range(8):
piece = board[file * 8 + rank]
if piece == "<empty>":
empty += 1
continue
if piece not in PIECES:
raise ValueError(f"{piece!r} is not a piece token")
row += (str(empty) if empty else "") + piece
empty = 0
ranks.append(row + (str(empty) if empty else ""))
turn = names[1 + N_SQUARES].removeprefix("turn:")
castling = names[2 + N_SQUARES].removeprefix("castle:")
ep_file = names[3 + N_SQUARES].removeprefix("ep:")
ep = "-" if ep_file == "none" else f"{ep_file}{'6' if turn == 'w' else '3'}"
return f"{'/'.join(ranks)} {turn} {castling} {ep}"
def square_name(index: int) -> str:
"""Name of the board square at position ``index`` of the 64-square block (0 = ``a1``)."""
return squares()[index]

src/rukh/models/squares.pylíneas 176-215 · p3

La inversa existe para poder comprobar la directa. El token al paso solo guarda la columna, porque la fila se deduce del turno ('6' if turn == 'w' else '3'). Y como el tramo de reloj es una compresión con pérdida, el test de abajo compara contra el FEN de cuatro campos.

El test que caza el tablero girado

tests/unit/test_squares.py
"""Tests for rukh.models.squares: the 69-token FEN scheme, its vocabulary and the round trip."""
from __future__ import annotations
import chess
import pytest
from rukh.models.squares import (
CLOCK_EDGES,
SQUARE_IDS,
SQUARE_TOKENS,
SQUARE_VOCAB,
SQUARE_VOCAB_SIZE,
castling_strings,
clock_bucket,
fen_to_tokens,
square_name,
tokens_to_fen,
tokens_to_strings,
vocab_hash,
)
from rukh.tokenize.uci_vocab import squares
pytestmark = pytest.mark.unit
# Pinned like the UCI vocabulary: an encoder is only meaningful with the vocabulary it saw.
SQUARE_VOCAB_SHA256 = "b0f530895c035c07423dfa246ad5e874c7406ffbaa1b2155601ebf09e89c98e9"
OPENING = ["e2e4", "c7c5", "g1f3", "d7d6", "d2d4", "c5d4", "f3d4", "g8f6", "b1c3", "a7a6"]
def played(moves: list[str]) -> chess.Board:
board = chess.Board()
for move in moves:
board.push_uci(move)
return board
def test_the_vocabulary_is_stable() -> None:
assert vocab_hash() == SQUARE_VOCAB_SHA256
assert len(SQUARE_VOCAB) == SQUARE_VOCAB_SIZE == 47
assert len(set(SQUARE_VOCAB)) == SQUARE_VOCAB_SIZE
assert SQUARE_VOCAB[:4] == ["<pad>", "<mask>", "<cls>", "<empty>"]
assert SQUARE_VOCAB[4:16] == list("PNBRQKpnbrqk")
assert len(castling_strings()) == 16
assert castling_strings()[0] == "-" and castling_strings()[-1] == "KQkq"

tests/unit/test_squares.pylíneas 1-45 · p3

El SHA-256 escrito a mano es la técnica del test_lock_guard.py de M0. Un cambio deliberado del vocabulario obliga a cambiar el hash y a decidir qué pasa con los checkpoints entrenados con el anterior, que es la conversación que hay que tener.

tests/unit/test_squares.py
def test_a_position_is_exactly_69_tokens() -> None:
board = played(OPENING)
assert len(fen_to_tokens(board.fen())) == SQUARE_TOKENS == 69
fen4 = " ".join(board.fen().split()[:4])
assert len(fen_to_tokens(fen4)) == SQUARE_TOKENS # the P1 four-field FEN is accepted too
def test_the_layout_is_cls_board_turn_castling_en_passant_clock() -> None:
names = tokens_to_strings(fen_to_tokens(chess.Board().fen()))
assert names[0] == "<cls>"
assert names[65:] == ["turn:w", "castle:KQkq", "ep:none", "clock:0"]
assert names[1] == "R" and names[2] == "P" and names[3] == "<empty>" # a1, a2, a3
def test_every_square_matches_python_chess() -> None:
board = chess.Board()
for move in OPENING:
board.push_uci(move)
names = tokens_to_strings(fen_to_tokens(board.fen()))[1:65]
for index, name in enumerate(names):
piece = board.piece_at(chess.square(index // 8, index % 8))
assert name == (piece.symbol() if piece else "<empty>"), square_name(index)

tests/unit/test_squares.pylíneas 48-69 · p3

El tercero caza la transposición de la única manera fiable: contra una implementación independiente. Comprueba las 64 casillas de diez posiciones contra python-chess; comprobarlas con otra función del mismo fichero solo demostraría que las dos se equivocan igual. La demo usa el mismo principio en la lección 10 para su tokenizador en TypeScript. El square_name(index) del assert hace que el fallo diga «g4» en vez de «índice 30».

tests/unit/test_squares.py
def test_the_64_squares_round_trip_through_the_four_field_fen() -> None:
board = chess.Board()
for move in [*OPENING, "f1e2", "c8d7", "e1g1"]:
board.push_uci(move)
fen4 = " ".join(board.fen().split()[:4])
assert tokens_to_fen(fen_to_tokens(fen4)) == fen4
assert tokens_to_fen(fen_to_tokens(board.fen())) == fen4
def test_castling_rights_are_distinguished() -> None:
board = played(OPENING)
fen4 = " ".join(board.fen().split()[:4])
fields = fen4.split()
variants = {
rights: fen_to_tokens(f"{fields[0]} {fields[1]} {rights} {fields[3]}")[66]
for rights in ("KQkq", "KQk", "Kq", "-")
}
assert len(set(variants.values())) == 4
assert variants["-"] == SQUARE_IDS["castle:-"]
# The order inside the field does not matter, the rights do.
assert fen_to_tokens(f"{fields[0]} {fields[1]} qK {fields[3]}")[66] == variants["Kq"]
with pytest.raises(ValueError, match="unsupported castling field"):
fen_to_tokens(f"{fields[0]} {fields[1]} AHah {fields[3]}")

tests/unit/test_squares.pylíneas 72-94 · p3

La lista de jugadas del primero termina en e1g1, un enroque corto: mueve dos piezas y quita derechos, así que es la jugada que más puede romper. El segundo comprueba que el empaquetado del enroque no pierde ningún bit y que el orden dentro del campo no importa.

tests/unit/test_squares.py
def test_en_passant_is_distinguished_and_does_not_disturb_the_board() -> None:
board = played(["e2e4", "a7a6", "e4e5", "d7d5"]) # d6 is a legal en-passant target
with_ep = fen_to_tokens(board.fen())
without = fen_to_tokens(board.fen().replace(" d6 ", " - "))
assert with_ep[67] == SQUARE_IDS["ep:d"] and without[67] == SQUARE_IDS["ep:none"]
assert with_ep[:67] == without[:67] # only the en-passant slot changes
assert tokens_to_fen(with_ep).split()[3] == "d6"
assert fen_to_tokens("8/8/8/8/8/8/8/8 b - c3")[67] == SQUARE_IDS["ep:c"]
assert tokens_to_fen(fen_to_tokens("8/8/8/8/8/8/8/8 b - c3")).split()[3] == "c3"
with pytest.raises(ValueError, match="en-passant square"):
fen_to_tokens("8/8/8/8/8/8/8/8 w - e4")
def test_the_halfmove_clock_is_bucketed() -> None:
assert CLOCK_EDGES == (6, 25, 50)
assert [clock_bucket(n) for n in (0, 5, 6, 24, 25, 49, 50, 120)] == [0, 0, 1, 1, 2, 2, 3, 3]
start = chess.Board().fen().split()
fen = " ".join(start[:4])
assert fen_to_tokens(f"{fen} 0 1")[68] == SQUARE_IDS["clock:0"]
assert fen_to_tokens(f"{fen} 30 40")[68] == SQUARE_IDS["clock:2"]
assert fen_to_tokens(fen)[68] == SQUARE_IDS["clock:0"] # a fen4 carries no counter
with pytest.raises(ValueError, match="cannot be negative"):
clock_bucket(-1)

tests/unit/test_squares.pylíneas 97-119 · p3

with_ep[:67] == without[:67] es una aserción de aislamiento: cambiar un campo del FEN solo puede mover su ranura. Sirve para cualquier formato de ancho fijo. Y el test del reloj recorre las fronteras (5 y 6, 24 y 25), porque un > donde debía haber un >= solo se nota ahí.

tests/unit/test_squares.py
def test_the_square_order_is_the_uci_vocabulary_order() -> None:
assert square_name(0) == "a1" and square_name(63) == "h8"
assert [square_name(i) for i in range(64)] == squares()
def test_a_malformed_fen_is_refused() -> None:
with pytest.raises(ValueError, match="at least 4 fields"):
fen_to_tokens("8/8/8/8/8/8/8/8 w")
with pytest.raises(ValueError, match="8 ranks"):
fen_to_tokens("8/8/8 w - -")
with pytest.raises(ValueError, match="covers 7 files"):
fen_to_tokens("7/8/8/8/8/8/8/8 w - -")
with pytest.raises(ValueError, match="unknown piece"):
fen_to_tokens("xxxxxxxx/8/8/8/8/8/8/8 w - -")
with pytest.raises(ValueError, match="side to move"):
fen_to_tokens("8/8/8/8/8/8/8/8 x - -")
with pytest.raises(ValueError, match="69 tokens"):
tokens_to_fen([1, 2, 3])

tests/unit/test_squares.pylíneas 122-139 · p3

El último convierte «un FEN mal formado lanza» en una promesa comprobada: seis maneras de escribir mal un FEN, cada una con su mensaje. Los match solo miran el trozo que identifica el caso, para no romperse cada vez que alguien mejora la redacción.

// Ejercicio 01Gira el tablero y mira qué tests sobreviven

En una copia del repositorio, cambia ids[file * 8 + rank] por ids[rank * 8 + file] en _placement_ids y ejecuta uv run pytest tests/unit/test_squares.py -q. ¿Cuáles de los diez tests fallan? Después deshazlo y quita el <cls> de build_square_vocab y de la lista que devuelve fen_to_tokens, para que la secuencia tenga 68 tokens: ¿qué se rompe y qué no?

// SoluciónVer la solución

Con la transposición girada fallan test_the_layout_is_cls_board_turn_castling_en_passant_clock (la casilla 1 ya no es la torre de a1) y test_every_square_matches_python_chess, que es el que lo dice bien: enseña la casilla concreta en la que discrepa de python-chess. Lo interesante es cuál no falla: test_the_64_squares_round_trip_through_the_four_field_fen pasa tranquilamente, porque tokens_to_fen usa la misma fórmula y las dos se equivocan igual. Una inversa consistente no demuestra que la directa sea correcta, solo que es invertible.

Sin <cls>, lo primero que rompe es test_the_vocabulary_is_stable, por el hash y por el tamaño; y test_a_position_is_exactly_69_tokens y todos los que indexan 65, 66, 67 o 68. Lo que no rompe es el modelo: EncoderConfig.seq seguiría diciendo 69 porque lee SQUARE_TOKENS, así que habría que cambiar esa constante también, y entonces la tabla de posiciones de cualquier checkpoint entrenado antes dejaría de encajar. Esa cadena es la razón de que el hash del vocabulario sea un test y no un comentario.

Qué has aprendido

Ponerle tres cabezas a un tronco exige una pérdida con pesos explícitos, las partes registradas por separado y una máscara para la tarea cuya etiqueta a veces no existe. Una representación de ancho fijo se defiende con un vocabulario construido y con hash, una inversa que existe para comprobar la directa y una comparación contra una implementación ajena. Te servirán para cualquier formato de entrada que diseñes.

Cómo se mide: uv run pytest tests/unit/test_squares.py -q pasa los diez tests, y el hash b0f5308… es el del vocabulario de 47 tokens. Los tests de las cabezas de test_heads.py pasan sin GPU y en menos de un segundo en cuanto exista el afinado de la lección 6.

Lo siguiente es el preentrenamiento: cómo se tapa el 15 % de las jugadas de una partida, por qué el sorteo son dos sorteos y qué comparte el bucle del encoder con el del decoder de M2.