// 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.
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
"""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 theexperiment mean something. If a linear probe on the pooled representation can say how good aposition is, whether the last move threw the game away and who is going to win, then therepresentation already contains those facts and masked move modeling put them there. A deephead 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 threetasks have different scales (an MSE in ``[0, 4]``, two cross entropies) and different amounts ofdata, 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 torchfrom torch import Tensor, nnfrom torch.nn import functional as F
from rukh.config import BaseConfigfrom rukh.models.encoder import PositionEncoder
HEADS = ("value", "blunder", "result")RESULT_CLASSES = 3El 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.
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.5Hereda 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.
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)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.
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), }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.
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, partsTres 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 elzero. Un lote sin filas etiquetadas daría una entropía cruzada sobre tensores vacíos, que esnan, y esenanllegaría por la suma a los gradientes de las tres cabezas y del tronco. Un solo lote desafortunado arruinaría la tirada; con elzero, 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.
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() ]Existe para que la lección 6 pueda escribir «descongela esto» en una línea.
Los tests de las cabezas
"""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 plimport pytestimport torch
from rukh.models import EncoderConfig, MultiHead, PositionEncoderfrom rukh.models.heads import HEADS, BlunderHead, HeadWeights, ResultHead, ValueHeadfrom rukh.models.squares import SQUARE_TOKENS, fen_to_tokensfrom rukh.train import HeadsConfig, freeze_encoder, label_curve, train_headsfrom rukh.train.checkpoint import load_checkpoint, save_checkpointfrom 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 -"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.
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]), }La mitad de las filas del lote no tienen etiqueta de error, para comprobar la máscara sin montar un dataset.
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_TOKENSEl primero caza la difusión silenciosa del squeeze y, de paso, que el tanh sigue puesto.
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 tooComprueba 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.
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)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.
"""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); thisone feeds it the position itself, so the two can be compared on the same heads. That comparisonis 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 theUCI 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 69it fixes; 69 is the binding number, so the four castling rights travel in one token with 16values (they are four bits of one fact, and this keeps every other component in a slot of itsown) and the position that frees up became ``<cls>``, which gives ``pool("cls")`` a real anchorin 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, soit 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 UCIvocabulary of P1."""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:0en 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.
from __future__ import annotations
import hashlibimport json
from rukh.tokenize.uci_vocab import FILES, squares
PIECES = "PNBRQKpnbrqk"CASTLING_ORDER = "KQkq"N_SQUARES = 64SQUARE_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."""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.
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)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.
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 idsLa 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.
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}"]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).
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)}"], ]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.
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]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 for rukh.models.squares: the 69-token FEN scheme, its vocabulary and the round trip."""
from __future__ import annotations
import chessimport 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"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.
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)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».
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]}")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.
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)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í.
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])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.