rukh · lab

// M3 · lección 03

El encoder a mano: el mismo bloque con `causal=False`

`models/encoder.py` entero: la máscara de claves del relleno, la guarda que impide un NaN y por qué se apaga bajo el exportador, la cabeza de jugada tapada con su `-100`, el pooling que ignora el relleno, y el test que exige lo contrario que el de M2.

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

Lección 3 de 11 del módulo «El encoder». Viene de «Las capas compartidas» y sigue en «Cabezas y casillas».

Qué vas a construir

src/rukh/models/encoder.py y su fichero de tests. Con layers.py en su sitio, el modelo entero cabe en una clase: el bucle de bloques con causal=False, la máscara de claves que apaga el relleno, la cabeza que predice la jugada tapada y el poolingPoolingReducir los T vectores que devuelve un encoder a uno solo que represente la secuencia entera. Rukh implementa los dos clásicos: cls toma el vector de la primera posición (<bos> en la entrada de jugadas, <cls> en la de casillas) y mean promedia solo los tokens reales, nunca el relleno. Ese único vector es lo que leen las tres cabezas de M3 y lo que se guarda como embedding de posición. que reduce la secuencia a un vector.

El contrato, escrito antes que el código

src/rukh/models/encoder.py
"""``PositionEncoder``: a hand-written bidirectional transformer over positions.
The brother of ``MoveDecoder``, and deliberately built from the same ``rukh.models.layers``
blocks: the whole difference is the attention mask. The decoder answers "what comes next?", so
a position may only look left; the encoder answers "what is this position?", so every token
sees every other one, including the ones after it. That is why the test that proves the decoder
right (a future token never changes a past logit) has to fail here, and ``test_encoder`` asserts
the opposite.
Two input schemes share the class (``EncoderConfig.input``):
``moves``
the game so far in the P1 UCI vocabulary, the very tokens the decoder was trained on, so
the masked-move head can be tied to the embedding and a pretrained decoder's intuition is
directly comparable.
``squares``
the 69 tokens of ``rukh.models.squares``, the position itself rather than its history.
Padding is expressed as a key mask, not as a special attention: ``attention_mask`` is ``True``
on real tokens and ``False`` on ``<pad>``, and the padded keys are switched off for every
query. The rows of padded queries are still computed (nobody reads them) and at least one real
token per row is required, otherwise softmax would see an entirely masked row and return NaN.
"""
from __future__ import annotations
import math
from typing import Literal
import torch
from torch import Tensor, nn
from torch.nn import functional as F
from rukh.models import layers
from rukh.models.config import EncoderConfig
from rukh.models.layers import rope_tables

src/rukh/models/encoder.pylíneas 1-36 · p3

El último párrafo contiene la decisión: el relleno se expresa como máscara de claves. Cada posición puede consultar a todas menos a las de relleno. Enmascarar también las consultas, para no calcular las filas del relleno, suena más eficiente y es lo que produce NaN: una fila con todas sus claves apagadas le da al softmax una fila entera de -inf, y el softmax de eso es indefinido. Rukh calcula esas filas, no las mira nunca y exige que cada secuencia tenga al menos un token real.

Dos constantes, y una de ellas es la lección

src/rukh/models/encoder.py
PAD_ID = 0
"""``<pad>`` is id 0 in both schemes."""
MMM_IGNORE_INDEX = -100
"""Label of a position the masked-move loss must skip.
Not the decoder's ``0``. ``0`` is ``<pad>`` in **both** schemes, so using it to mean "nothing to
predict" would overload one id with two jobs: the loss could no longer tell a position it must
skip from a position where the right answer happens to be ``<pad>``. The decoder gets away with
it because its target is the next token of a packed stream and ``<pad>`` is never a target there;
here the distinction has to be explicit, and ``-100`` is outside every vocabulary, so "not
predicted" and "predict ``<pad>``" stay different things.
"""

src/rukh/models/encoder.pylíneas 38-50 · p3

El docstring de MMM_IGNORE_INDEX es el porqué que la lección 1 anticipó, escrito donde vive la constante: aquí <pad> puede ser un objetivo legítimo, así que «no puntúes esto» tiene que ser un número que no sea el id de nada.

La función que existe solo para el exportador

src/rukh/models/encoder.py
def _tracing() -> bool:
"""``True`` while ``torch.export`` or ``torch.compile`` is capturing a graph.
The padding check below reads a tensor to decide whether to raise, and a data-dependent
branch like that is precisely what a graph capture cannot represent. Under ``torch.export``
it would send the exporter back to the deprecated TorchScript tracer (D-027); under
``torch.compile`` it is worse in a quieter way — Dynamo cannot prove the condition, so it
graph-breaks around it on **every** masked step of the MMM loop, which is a synchronisation
point plus two half-graphs per step for a check that has already passed. Eager behaviour is
unchanged: outside a capture both calls are ``False`` and the ValueError is raised as before.
"""
return torch.compiler.is_exporting() or torch.compiler.is_compiling()

src/rukh/models/encoder.pylíneas 53-64 · p3

Casi toda la función es su docstring, y con razón: es el tipo de código que alguien borra por «limpieza» en un año si no está escrito por qué está.

Una comprobación que lee un tensor para decidir si lanza —aquí, if not mask.any(...)— es una bifurcación que depende de los datos, y una captura de grafo no puede representarla: el grafo tiene que ser el mismo para cualquier entrada. Bajo torch.export, el exportador vuelve al tracer antiguo de TorchScript, el camino que M2 dejó cerrado. Bajo torch.compile, Dynamo corta el grafo alrededor de la condición en cada paso, con una sincronización entre CPU y GPU cada vez. La comprobación protege de un NaN real, así que no se quita: se apaga solo mientras alguien captura, y fuera de una captura todo se comporta como siempre.

El constructor: causal=False y un atado condicional

src/rukh/models/encoder.py
class PositionEncoder(nn.Module):
"""Bidirectional transformer encoder: ``forward`` returns the hidden states ``(B, T, d)``."""
def __init__(self, cfg: EncoderConfig | None = None) -> None:
super().__init__()
self.cfg = cfg or EncoderConfig()
self.tokens = nn.Embedding(self.cfg.tokens, self.cfg.d_model)
self.positions = (
nn.Embedding(self.cfg.seq, self.cfg.d_model) if self.cfg.pos == "learned" else None
)
self.drop = nn.Dropout(self.cfg.dropout)
self.blocks = nn.ModuleList(
layers.Block(
self.cfg.d_model, self.cfg.n_head, self.cfg.ff, self.cfg.dropout, causal=False
)
for _ in range(self.cfg.n_layer)
)
self.ln_f = nn.LayerNorm(self.cfg.d_model)
self.mlm_head = nn.Linear(self.cfg.d_model, self.cfg.tokens, bias=False)
if self.cfg.tie_embeddings and self.cfg.input == "moves":
# Tied only for ``moves``: it is the P1 vocabulary, where the embedding and the head
# describe the same 2 030 moves. The 47 square tokens are too few to be worth tying.
self.mlm_head.weight = self.tokens.weight
if self.cfg.pos == "rope":
cos, sin = rope_tables(self.cfg.seq, self.cfg.head_dim, torch.device("cpu"))
self.register_buffer("rope_cos", cos, persistent=False)
self.register_buffer("rope_sin", sin, persistent=False)
self.apply(layers.init_weights)
scale = 0.02 / math.sqrt(2 * self.cfg.n_layer)
for name, param in self.named_parameters():
if name.endswith("proj.weight"):
nn.init.normal_(param, mean=0.0, std=scale)

src/rukh/models/encoder.pylíneas 67-98 · p3

Léelo al lado del __init__ de MoveDecoder de M2: es la misma secuencia. Cambian causal=False en los bloques y el nombre de la cabeza, mlm_head, porque predice una jugada tapada y no la siguiente. Tres detalles son decisiones:

  • self.cfg.tokens y self.cfg.seq resuelven el esquema, así que el constructor no pregunta por cfg.input para decidir un tamaño. El único if que mira el esquema es el del atado.
  • El atado es condicional. Solo tiene sentido si la entrada y la salida son el mismo vocabulario: con 2 030 jugadas ahorra 779 520 parámetros, el 5 % del modelo, e impone que un embeddingEmbeddingTabla que asigna un vector aprendido a cada id del vocabulario, y por extensión ese vector. En rukh-small la tabla es de 2 030 × 512: cada jugada UCI entra en el modelo como un punto en un espacio de 512 dimensiones, aprendido a la vez que el resto de la red. signifique lo mismo al entrar y al salir. Con 47 tokens de casillas se ahorrarían 18 048 y se obligaría al embedding de entrada a hacer de clasificador sin ninguna razón.
  • La inicialización escalada de las proyecciones residuales, 0.02 / sqrt(2 · n_layer), es la de GPT-2 y va después de self.apply(layers.init_weights); al revés, la general la borraría. Cada bloque suma dos veces al residuo, así que con ocho capas hay dieciséis sumas y la varianza crecería con la profundidad; dividir por sqrt(2 · n_layer) lo compensa.

Contar parámetros y leer el relleno

src/rukh/models/encoder.py
def num_params(self, non_embedding: bool = True) -> int:
"""Parameter count; ``non_embedding`` drops the learned position table (nanoGPT rule)."""
total = sum(p.numel() for p in self.parameters())
if non_embedding and self.positions is not None:
total -= self.positions.weight.numel()
return total
@staticmethod
def padding_mask(idx: Tensor) -> Tensor:
"""``True`` where ``idx`` is a real token, ``False`` on ``<pad>``; shape ``(B, T)``."""
return idx != PAD_ID

src/rukh/models/encoder.pylíneas 100-110 · p3

num_params es el de M2, y con non_embedding=False da los 15 052 800 de la lección 1. Ojo con la convención de nanoGPT: resta la tabla de posiciones, no la de tokens, aunque el nombre diga «embedding». Cuando se publica un tamaño se publica el total.

padding_mask es un != con nombre: sin él, idx != 0 aparecería suelto en cinco sitios y el día que <pad> dejara de ser el 0 habría que encontrarlos todos.

La guarda que impide el NaN

src/rukh/models/encoder.py
@staticmethod
def _key_mask(attention_mask: Tensor | None) -> Tensor | None:
"""``(B, T)`` into the ``(B, 1, 1, T)`` boolean key mask SDPA expects."""
if attention_mask is None:
return None
mask = attention_mask.bool()
if mask.dim() != 2:
raise ValueError(f"attention_mask must be (B, T), got {tuple(attention_mask.shape)}")
if not _tracing() and not bool(mask.any(dim=-1).all()):
# An entirely masked row would make softmax return NaN, so it is refused here rather
# than debugged three layers down. The check reads a tensor, which is exactly what
# `torch.export` cannot trace (a data-dependent guard), and skipping it under the
# exporter is what keeps the encoder on the modern exporter instead of the
# deprecated tracer: the exported graph is a pure function of its input either way.
raise ValueError("every sequence needs at least one unmasked token")
return mask[:, None, None, :]

src/rukh/models/encoder.pylíneas 112-127 · p3

Hace tres cosas. mask[:, None, None, :] convierte (B, T) en (B, 1, 1, T): los dos ejes de tamaño uno son la cabeza y la consulta, y se difunden, así que la misma máscara vale para las seis cabezas y las doscientas consultas sin materializar veintitrés millones de booleanos por capa.

La comprobación de forma existe porque un attention_mask de (B, 1, T), lo que devuelve más de una librería, se difundiría sin protestar contra el eje de las cabezas y enmascararía otra cosa.

Y la guarda del NaN: mask.any(dim=-1).all() exige que a cada secuencia le quede algún token real. Sin ella, el NaN del softmax se propagaría en silencio y aparecería en el loss ocho capas después, cuando ya no hay forma de saber de dónde salió. Solo se salta bajo una captura de grafo.

El forward, que es el del decoder con un argumento más

src/rukh/models/encoder.py
def forward(self, idx: Tensor, attention_mask: Tensor | None = None) -> Tensor:
"""Hidden states ``(B, T, d_model)``; ``attention_mask`` is ``True`` on real tokens."""
seq = idx.shape[-1]
if seq > self.cfg.seq:
raise ValueError(f"sequence of {seq} tokens is longer than block {self.cfg.seq}")
key_mask = self._key_mask(attention_mask)
x = self.tokens(idx)
cos = sin = None
if self.positions is not None:
steps = torch.arange(seq, device=idx.device)
x = x + self.positions(steps)
else:
cos, sin = self.rope_cos, self.rope_sin
x = self.drop(x)
for block in self.blocks:
x = block(x, cos, sin, key_mask)
return self.ln_f(x)

src/rukh/models/encoder.pylíneas 129-145 · p3

La diferencia con el forward de MoveDecoder está en lo que devuelve: estados ocultos, no logits. El decoder termina en lm_head porque solo tiene una pregunta; el encoder devuelve (B, T, d_model) y deja que quien llame decida: la cabeza de jugada tapada, el pooling, las cabezas supervisadas o los embeddings que se exportan en la lección 9.

La comprobación de longitud va antes de todo porque el fallo natural sería peor: con posiciones aprendidas, un error de índice dentro de nn.Embedding que no menciona la longitud; con RoPE, ni siquiera un error, porque cos[:t] se quedaría corto en silencio.

La cabeza de jugada tapada

src/rukh/models/encoder.py
def masked_lm(
self, idx: Tensor, labels: Tensor | None = None, attention_mask: Tensor | None = None
) -> tuple[Tensor, Tensor | None]:
"""Masked-move logits ``(B, T, V)`` and, with ``labels``, the scalar loss.
``labels`` is ``MMM_IGNORE_INDEX`` wherever nothing was masked; see ``rukh.train.mmm``.
"""
logits = self.mlm_head(self(idx, attention_mask))
loss = None
if labels is not None:
loss = F.cross_entropy(
logits.reshape(-1, logits.shape[-1]),
labels.reshape(-1),
ignore_index=MMM_IGNORE_INDEX,
)
return logits, loss

src/rukh/models/encoder.pylíneas 147-162 · p3

Es la pérdida de la MLMMLM (masked language modeling)Objetivo de preentrenamiento de BERT: se esconde una parte de los tokens y el modelo, que ve la secuencia por los dos lados, tiene que reconstruirlos. En Rukh se llama masked move modeling porque el token es una jugada: se tapa el 15 % de las jugadas de la partida y de esas el 80 % se sustituye por <mask>, el 10 % por otra jugada al azar y el 10 % se deja tal cual. Los tokens de control (<bos>, Elo, resultado, <eos>) nunca se tapan: son la condición, no la señal. y nada más. La sutileza está en ignore_index: con reduction="mean", la entropía cruzada divide por el número de posiciones no ignoradas, no por N, así que un lote donde el sorteo tapó pocas jugadas no produce un número artificialmente pequeño. El precio es un caso degenerado: si ninguna posición está tapada, el denominador es cero y la pérdida sale nan. Es aritmética, y por eso el bucle de la lección 5 comprueba torch.isfinite(loss) antes de acumular.

Devolver (logits, loss) permite que la validación calcule la exactitud top-1 con los mismos logits, sin una segunda pasada.

El pooling, y por qué el clamp

src/rukh/models/encoder.py
def pool(
self,
idx: Tensor,
how: Literal["cls", "mean"] = "mean",
attention_mask: Tensor | None = None,
) -> Tensor:
"""One vector per sequence, ``(B, d_model)``: the position's representation.
``cls`` takes position 0 (``<bos>`` for ``moves``, ``<cls>`` for ``squares``) and
``mean`` averages the real tokens only. Without an explicit ``attention_mask`` the
padding is read off ``idx`` itself, so ``pool(idx, "mean")`` never averages ``<pad>``
into the representation, which is the whole point of the operation.
"""
if how not in ("cls", "mean"):
raise ValueError(f"how must be 'cls' or 'mean', got {how!r}")
mask = self.padding_mask(idx) if attention_mask is None else attention_mask.bool()
hidden = self(idx, mask)
if how == "cls":
return hidden[:, 0]
weights = mask.unsqueeze(-1).to(hidden.dtype)
return (hidden * weights).sum(dim=1) / weights.sum(dim=1).clamp(min=1.0)

src/rukh/models/encoder.pylíneas 164-184 · p3

Es la operación que convierte un encoder en un extractor de representaciones, con tres decisiones dentro:

  • La máscara se deriva de idx cuando no se la pasan, así que el caso correcto es el caso por defecto: pool(idx) promedia los tokens reales y nunca el relleno.
  • La media se escribe como suma ponderada partida por el peso total, no como hidden[mask].mean(), que aplanaría el lote entero y perdería a qué secuencia pertenece cada vector. Multiplicar por una máscara y sumar por el eje del tiempo lo hace por lote, sin ramas.
  • clamp(min=1.0) en el denominador, porque pool acepta una máscara explícita que _key_mask no validó, y una división por cero en PyTorch no lanza: da inf o nan.

El if how not in (…) tampoco es burocracia: how viene de un YAML, y un pooling: media caería en silencio en la rama de la media si la comprobación fuera if how == "cls": … else: media.

El test que tiene que decir lo contrario que el de M2

tests/unit/test_encoder.py
"""Tests for rukh.models.encoder: shapes, bidirectionality, padding, pooling and the preset."""
from __future__ import annotations
import pytest
import torch
from rukh.models import EncoderConfig, PositionEncoder
from rukh.models.encoder import MMM_IGNORE_INDEX
from rukh.models.squares import SQUARE_TOKENS, SQUARE_VOCAB_SIZE, fen_to_tokens
pytestmark = pytest.mark.unit
TOY = EncoderConfig(vocab_size=64, n_layer=2, n_head=4, d_model=32, block=16, dropout=0.0)
START = "rnbqkbnr/pppppppp/8/8/8/8/PPPPPPPP/RNBQKBNR w KQkq -"
def toy_model(seed: int = 0, **overrides: object) -> PositionEncoder:
"""A deterministic two-layer encoder in eval mode."""
torch.manual_seed(seed)
return PositionEncoder(TOY.model_copy(update=overrides)).eval()
def toy_batch(batch: int = 2, seq: int = 8, seed: int = 1) -> torch.Tensor:
generator = torch.Generator().manual_seed(seed)
return torch.randint(1, TOY.vocab_size, (batch, seq), generator=generator)

tests/unit/test_encoder.pylíneas 1-26 · p3

El modelo de juguete tiene dos capas y 32 dimensiones porque un test que tarde un segundo se ejecuta y uno que tarde un minuto no. dropout=0.0 y .eval() a la vez, porque la mitad de estos tests comparan dos pasadas y cualquier aleatoriedad los haría fallar de vez en cuando. Y torch.randint(1, …) empieza en 1 porque el 0 es <pad>.

tests/unit/test_encoder.py
def test_forward_returns_one_hidden_state_per_token() -> None:
hidden = toy_model()(toy_batch())
assert hidden.shape == (2, 8, TOY.d_model)
@pytest.mark.parametrize("pos", ["learned", "rope"])
def test_a_later_token_does_change_the_earlier_outputs(pos: str) -> None:
"""The inverse of the decoder's causality test: this model is bidirectional on purpose."""
model = toy_model(pos=pos)
idx = toy_batch(seq=12)
cut = 5
changed = idx.clone()
changed[:, cut + 1 :] = (changed[:, cut + 1 :] + 7) % (TOY.vocab_size - 1) + 1
assert not torch.equal(idx, changed)
with torch.no_grad():
base, other = model(idx), model(changed)
assert not torch.allclose(base[:, : cut + 1], other[:, : cut + 1], atol=1e-5)

tests/unit/test_encoder.pylíneas 29-45 · p3

Este es el test del módulo. En M2 el mismo experimento terminaba en assert torch.allclose(…): cambiar un token futuro no podía mover ni un decimal de un logit pasado. Aquí el assert lleva un not delante.

(x + 7) % (vocab_size - 1) + 1 cambia la cola sin producir nunca un 0, para que el resultado no pueda explicarse por el relleno. Y el assert not torch.equal(idx, changed) comprueba que la transformación cambió algo: sin él, el test podría pasar por el motivo equivocado, el fallo clásico de una aserción negativa. El parametrize sobre pos cubre RoPE, que mete la posición dentro de la atención y es otro camino de código.

tests/unit/test_encoder.py
def test_padding_does_not_leak_into_the_real_tokens() -> None:
model = toy_model()
short = toy_batch(seq=6, seed=2)
padded = torch.cat([short, torch.zeros((2, 4), dtype=torch.long)], dim=1)
with torch.no_grad():
plain = model(short, model.padding_mask(short))
with_padding = model(padded, model.padding_mask(padded))
assert torch.allclose(plain, with_padding[:, :6], atol=1e-5)
# Without the mask the padding is just another token and does change the answer.
with torch.no_grad():
unmasked = model(padded)
assert not torch.allclose(plain, unmasked[:, :6], atol=1e-5)
def test_an_entirely_masked_sequence_is_refused() -> None:
model = toy_model()
idx = torch.zeros((2, 6), dtype=torch.long)
with pytest.raises(ValueError, match="at least one unmasked token"):
model(idx, model.padding_mask(idx))
with pytest.raises(ValueError, match=r"\(B, T\)"):
model(toy_batch(), torch.ones((2, 1, 8), dtype=torch.bool))

tests/unit/test_encoder.pylíneas 48-68 · p3

El primero prueba la máscara por los dos lados: con máscara la respuesta no depende del relleno, y sin máscara sí. La segunda mitad es la que impide que el test pase por casualidad. Es un patrón que vale la pena copiar: un test de un mecanismo comprueba que sin el mecanismo el resultado cambia. Si no, lo único que has demostrado es que el modelo no mira la entrada.

El segundo caza las dos formas de llamar mal a _key_mask. La barra invertida de r"\(B, T\)" está porque match es una expresión regular.

tests/unit/test_encoder.py
def test_mean_pooling_ignores_the_padding() -> None:
model = toy_model()
short = toy_batch(seq=6, seed=3)
padded = torch.cat([short, torch.zeros((2, 5), dtype=torch.long)], dim=1)
with torch.no_grad():
assert torch.allclose(model.pool(short, "mean"), model.pool(padded, "mean"), atol=1e-5)
assert torch.allclose(model.pool(short, "cls"), model.pool(padded, "cls"), atol=1e-5)
hidden = model(short, model.padding_mask(short))
assert torch.allclose(model.pool(short, "mean"), hidden.mean(dim=1), atol=1e-5)
assert torch.allclose(model.pool(short, "cls"), hidden[:, 0], atol=1e-6)
assert model.pool(short).shape == (2, TOY.d_model) # mean is the default
with pytest.raises(ValueError, match="cls"):
model.pool(short, "first") # type: ignore[arg-type]

tests/unit/test_encoder.pylíneas 71-83 · p3

La propiedad que se fija es la de la lección 1: la representación de una secuencia no puede depender de con quién le toque viajar en el lote. Las dos siguientes aserciones comprueban que la suma ponderada da lo mismo que un mean normal sin relleno, que es como se caza una errata en el eje.

tests/unit/test_encoder.py
def test_masked_lm_scores_only_the_masked_positions() -> None:
model = toy_model()
idx = toy_batch(seq=8, seed=4)
labels = torch.full_like(idx, MMM_IGNORE_INDEX)
labels[:, 2] = idx[:, 2]
logits, loss = model.masked_lm(idx, labels)
assert logits.shape == (2, 8, TOY.vocab_size)
assert loss is not None and loss.ndim == 0 and torch.isfinite(loss)
assert model.masked_lm(idx)[1] is None
# Moving a label to another position changes the loss; ignoring everything gives no loss.
elsewhere = torch.full_like(idx, MMM_IGNORE_INDEX)
elsewhere[:, 5] = idx[:, 5]
assert not torch.allclose(loss, model.masked_lm(idx, elsewhere)[1], atol=1e-6)
assert torch.isnan(model.masked_lm(idx, torch.full_like(idx, MMM_IGNORE_INDEX))[1])
# ``<pad>`` (0) is a legitimate label here, unlike in the decoder.
pad_label = torch.full_like(idx, MMM_IGNORE_INDEX)
pad_label[:, 1] = 0
assert torch.isfinite(model.masked_lm(idx, pad_label)[1])

tests/unit/test_encoder.pylíneas 86-104 · p3

Las dos últimas aserciones son el contrato de MMM_IGNORE_INDEX en código. torch.isnan(…) con todas las etiquetas ignoradas documenta el nan en vez de esconderlo, y rompe el test si alguien decide devolver 0.0 en ese caso, un cero que se sumaría a la media de la época sin que nadie lo note. La de pad_label dice que una etiqueta 0 se puntúa; en el decoder de M2, con ignore_index=0, habría desaparecido de la pérdida sin decir nada.

tests/unit/test_encoder.py
def test_the_masked_move_head_is_tied_only_for_moves() -> None:
moves = toy_model()
assert moves.mlm_head.weight is moves.tokens.weight
untied = toy_model(tie_embeddings=False)
assert untied.mlm_head.weight is not untied.tokens.weight
torch.manual_seed(0)
squares = PositionEncoder(EncoderConfig(input="squares", n_layer=2, n_head=4, d_model=32))
assert squares.mlm_head.weight is not squares.tokens.weight
assert squares.tokens.num_embeddings == SQUARE_VOCAB_SIZE
def test_the_squares_scheme_has_69_fixed_positions() -> None:
cfg = EncoderConfig(input="squares", n_layer=2, n_head=4, d_model=32, dropout=0.0)
assert cfg.seq == SQUARE_TOKENS and cfg.tokens == SQUARE_VOCAB_SIZE
torch.manual_seed(0)
model = PositionEncoder(cfg).eval()
assert model.positions is not None and model.positions.num_embeddings == SQUARE_TOKENS
idx = torch.tensor([fen_to_tokens(START), fen_to_tokens(f"{START} 30 40")])
with torch.no_grad():
assert model(idx).shape == (2, SQUARE_TOKENS, 32)
assert model.pool(idx, "cls").shape == (2, 32)
with pytest.raises(ValueError, match="longer than block"):
model(torch.zeros((1, SQUARE_TOKENS + 1), dtype=torch.long))

tests/unit/test_encoder.pylíneas 107-129 · p3

El primero comprueba el atado con is: pregunta si son el mismo tensor. Con torch.equal bastaría con tener los mismos números, y dos tensores inicializados con la misma semilla los tienen sin estar atados. El segundo falla si alguien construye la tabla de posiciones con cfg.block en lugar de cfg.seq.

tests/unit/test_encoder.py
def test_same_seed_gives_the_same_weights_and_hidden_states() -> None:
idx = toy_batch()
with torch.no_grad():
a, b, c = toy_model(seed=3)(idx), toy_model(seed=3)(idx), toy_model(seed=4)(idx)
assert torch.equal(a, b)
assert not torch.allclose(a, c, atol=1e-4)
def test_preset_size(capsys: pytest.CaptureFixture[str]) -> None:
counts = {
scheme: PositionEncoder(EncoderConfig(input=scheme)).num_params(non_embedding=False)
for scheme in ("moves", "squares")
}
with capsys.disabled():
print()
for scheme, total in counts.items():
print(f"encoder preset ({scheme:<7}) {total:>12,} parameters")
# 8 layers of d=384 are 14,195,712 parameters, plus 779,520 of embedding (2 030 moves) and
# 76,800 of position table: 15,052,800. The plan's "18M-25M" band was an estimate its own
# preset (n_layer=8, n_head=6, d_model=384) cannot reach; the preset is what is binding.
assert 14_000_000 <= counts["moves"] <= 16_000_000
assert counts["squares"] < counts["moves"] # 47 square tokens against 2 030 moves

tests/unit/test_encoder.pylíneas 132-153 · p3

El de la semilla tiene las dos mitades otra vez: sin la segunda, un modelo que devolviera ceros pasaría. El del tamaño comprueba una banda y no un número exacto, para que un cambio legítimo de la configuración no rompa un test que no trata de eso; capsys.disabled() imprime las dos cuentas aunque pytest capture la salida.

tests/unit/test_encoder.py
def test_config_rejects_unknown_keys_and_bad_shapes() -> None:
with pytest.raises(ValueError):
EncoderConfig(n_layers=3) # type: ignore[call-arg]
with pytest.raises(ValueError, match="not divisible"):
EncoderConfig(d_model=100, n_head=8)
with pytest.raises(ValueError, match="even head dimension"):
EncoderConfig(d_model=12, n_head=4, pos="rope")
with pytest.raises(ValueError, match="dropout"):
EncoderConfig(dropout=1.0)
with pytest.raises(ValueError, match="block"):
EncoderConfig(block=0)
assert EncoderConfig(d_model=64, n_head=4).ff == 256
assert EncoderConfig(d_model=64, n_head=4, d_ff=128).ff == 128
def test_a_few_steps_of_gradient_descent_reduce_the_loss() -> None:
torch.manual_seed(5)
model = PositionEncoder(TOY)
idx = toy_batch(seq=8, seed=6)
labels = torch.full_like(idx, MMM_IGNORE_INDEX)
labels[:, ::2] = idx[:, ::2]
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-2)
losses = []
for _ in range(10):
_, loss = model.masked_lm(idx, labels)
assert loss is not None
losses.append(loss.detach().item())
optimizer.zero_grad()
loss.backward()
optimizer.step()
assert losses[-1] < losses[0]

tests/unit/test_encoder.pylíneas 156-186 · p3

El primero recorre las validaciones de la lección 2; el n_layers=3 —con la ese— es el test de extra="forbid".

El último es el más barato y el que más veces salva. Diez pasos de AdamW sobre un lote fijo tienen que bajar la pérdida; si no bajan, algo estructural está roto: un detach de más, un no_grad mal puesto, una máscara que impide aprender. No comprueba que el modelo sea bueno, sino que puede aprender, antes de gastar dieciséis minutos de GPU. Conviene tenerlo en cualquier modelo que escribas.

// Ejercicio 01Rompe la bidireccionalidad y mira qué test se queja

En una copia del repositorio, cambia causal=False por causal=True en el __init__ de PositionEncoder y ejecuta uv run pytest tests/unit/test_encoder.py -q. Antes de mirar: ¿cuántos tests fallan, y cuáles? Después vuelve a dejarlo y quita el attn_mask=attn_mask de layers.SelfAttention.forward: ¿qué falla ahora?

// SoluciónVer la solución

Con causal=True falla test_a_later_token_does_change_the_earlier_outputs, las dos variantes del parametrize, y solo ese: todo lo demás del encoder es indiferente a la dirección.

Quitando el attn_mask de la llamada de atención fallan dos: test_padding_does_not_leak_into_the_real_tokens en su primera aserción —la respuesta pasa a depender de cuánto relleno viaje— y test_mean_pooling_ignores_the_padding en la suya, porque la media sigue dividiendo por los tokens reales pero los vectores que promedia ya están contaminados. Fíjate en que la segunda mitad del primer test, la que exige que sin máscara la respuesta sí cambie, seguiría pasando: por sí sola no demuestra nada.

Qué has aprendido

El modelo entero es el decoder de M2 más tres cosas que el decoder no necesitaba: una máscara de claves para el relleno, una etiqueta de «no puntúes esto» que no se confunde con ningún id, y una reducción de la secuencia a un vector que no promedia lo que no existe. Y un patrón de test para cualquier proyecto: para probar un mecanismo, comprueba también que sin él el resultado cambia.

Cómo se mide: uv run pytest tests/unit/test_encoder.py -q pasa los doce tests (trece casos, contando el parametrize), y entre ellos el que en M2 tenía que fallar. En tu repositorio se ejecutan al final de la lección siguiente, porque los tests y config.py importan rukh.models.squares y __init__.py importa las cabezas, y las dos cosas llegan allí.

Lo siguiente son las tres cabezas que leen ese vector y el traductor de FEN a 69 tokens, que es la segunda representación de entrada del experimento del módulo.