// 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.
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
"""``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?", soa position may only look left; the encoder answers "what is this position?", so every tokensees every other one, including the ones after it. That is why the test that proves the decoderright (a future token never changes a past logit) has to fail here, and ``test_encoder`` assertsthe 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 everyquery. The rows of padded queries are still computed (nobody reads them) and at least one realtoken per row is required, otherwise softmax would see an entirely masked row and return NaN."""
from __future__ import annotations
import mathfrom typing import Literal
import torchfrom torch import Tensor, nnfrom torch.nn import functional as F
from rukh.models import layersfrom rukh.models.config import EncoderConfigfrom rukh.models.layers import rope_tablesEl ú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
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 topredict" would overload one id with two jobs: the loss could no longer tell a position it mustskip from a position where the right answer happens to be ``<pad>``. The decoder gets away withit 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 "notpredicted" and "predict ``<pad>``" stay different things."""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
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()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
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)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.tokensyself.cfg.seqresuelven el esquema, así que el constructor no pregunta porcfg.inputpara decidir un tamaño. El únicoifque 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-smallla 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 deself.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 porsqrt(2 · n_layer)lo compensa.
Contar parámetros y leer el relleno
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_IDnum_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
@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, :]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
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)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
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, lossEs 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
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)Es la operación que convierte un encoder en un extractor de representaciones, con tres decisiones dentro:
- La máscara se deriva de
idxcuando 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, porquepoolacepta una máscara explícita que_key_maskno validó, y una división por cero en PyTorch no lanza: dainfonan.
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 for rukh.models.encoder: shapes, bidirectionality, padding, pooling and the preset."""
from __future__ import annotations
import pytestimport torch
from rukh.models import EncoderConfig, PositionEncoderfrom rukh.models.encoder import MMM_IGNORE_INDEXfrom 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)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>.
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)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.
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))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.
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]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.
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])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.
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))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.
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 movesEl 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.
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]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.