rukh · lab

// M2 · lección 02

El decoder a mano: config, bloques y el modelo

Los tres ficheros de `src/rukh/models/` enteros y línea a línea —la configuración que valida su propia forma, el decoder de 177 líneas con su atención fusionada y su inicialización escalada, y los reexportes— más los quince tests que convierten «funciona» en «sigue funcionando».

  • tiempo de trabajo110 min
  • nivel base
  • actualizado el22 de septiembre de 2026

Qué vas a construir

El modelo entero: tres ficheros, 257 líneas contando el __init__.py, y un fichero de tests de 159 más. Es poco, y ese es el argumento de escribirlo a mano en vez de importar nn.TransformerDecoder: cabe en una sentada, no hay nada que no puedas leer, y cuando en M3 el encoder sea «lo mismo sin la máscara causal» sabrás exactamente qué frase estás cambiando.

Todo lo que sigue está aquí entero y literal, en la etiqueta p2. Cada bloque lleva debajo su enlace a GitHub con las líneas exactas, y pnpm verify:code comprueba que sigue siendo el fichero. Si copias los bloques en orden, tienes src/rukh/models/.

config.py: donde cada decisión se convierte en un número

Es el fichero donde acaba todo lo que la lección anterior argumenta. Sin él, nada puede llamar a preset("small") ni leer cfg.ff.

src/rukh/models/config.py
"""Configuration of the hand-written GPT decoder and the three course presets."""
from __future__ import annotations
from typing import Literal
from pydantic import model_validator
from rukh.config import BaseConfig

src/rukh/models/config.pylíneas 1-9 · p2

BaseConfig es el modelo de pydantic con extra="forbid" que se escribió en M0: una clave desconocida en un YAML es un error, no un valor ignorado. Aquí importa más que en ninguna otra config del proyecto, porque un n_layers: 16 con la ese de más produciría un modelo de doce capas que nadie sabría que no es el que pidió.

src/rukh/models/config.py
class DecoderConfig(BaseConfig):
"""Shape of a ``MoveDecoder``.
``d_ff`` defaults to ``4 * d_model`` and ``pos`` picks learned positional embeddings or
rotary embeddings (RoPE) applied to the queries and keys of every head.
"""
vocab_size: int = 2030
n_layer: int = 12
n_head: int = 8
d_model: int = 512
d_ff: int | None = None # None -> 4 * d_model
block: int = 200
dropout: float = 0.0
pos: Literal["learned", "rope"] = "learned"
tie_embeddings: bool = True

src/rukh/models/config.pylíneas 12-27 · p2

Trece líneas y ocho decisiones. Las que no son evidentes:

  • d_ff: int | None = None en vez de 2048. El ancho del MLP casi siempre es 4 * d_model, y escribirlo como un número obligaría a cambiar dos campos para cambiar la talla del modelo. None significa «el que toca», y la propiedad ff de más abajo lo resuelve.
  • dropout: float = 0.0. Cero en el decoder, y no por descuido: con 240 millones de tokens y 39 millones de parámetros no hay escasez de datos contra la que regularizar, así que el dropout solo frenaría el ajuste. El encoder de M3 lo enciende, porque va a ver muchísimas menos etiquetas en su afinado.
  • pos: Literal["learned", "rope"] = "learned". Aprendidas por defecto, por el argumento de la lección anterior: con el contexto fijo en 200 ninguna partida es más larga que la tabla, así que extrapolar no compra nada. RoPERoPECodificación de posición rotatoria (rotary position embedding): en vez de sumar un vector de posición, rota por pares las dimensiones de las consultas y las claves con un ángulo proporcional a la posición. El producto escalar entre dos posiciones depende entonces solo de su distancia, y no hace falta una tabla de posiciones. En Rukh es la alternativa a las posiciones aprendidas (`pos: rope`). queda detrás de una bandera —y con su código escrito y probado— para que el día que alguien entrene las dos, la decisión cambie por una medida y no por una opinión.
  • block: int = 200 es el contexto, y aparece en cuatro sitios más del módulo: el recorte del prompt, el límite de plies de una partida, la longitud que traza el exportador y los metadatos del .onnx. Vive aquí porque es una propiedad del modelo, no del muestreador.
src/rukh/models/config.py
@model_validator(mode="after")
def _check(self) -> DecoderConfig:
if self.d_model % self.n_head:
raise ValueError(f"d_model={self.d_model} is not divisible by n_head={self.n_head}")
if min(self.vocab_size, self.n_layer, self.n_head, self.d_model, self.block) < 1:
raise ValueError("vocab_size, n_layer, n_head, d_model and block must be positive")
if not 0.0 <= self.dropout < 1.0:
raise ValueError(f"dropout must be in [0, 1), got {self.dropout}")
if self.d_ff is not None and self.d_ff < 1:
raise ValueError(f"d_ff must be positive, got {self.d_ff}")
if self.pos == "rope" and (self.d_model // self.n_head) % 2:
raise ValueError("rope needs an even head dimension")
return self

src/rukh/models/config.pylíneas 29-41 · p2

Un validador de pydantic que corre después de construir el objeto (mode="after"), así que ve los campos ya con su tipo. Las tres comprobaciones son las tres formas de pedir un modelo imposible: una anchura que no se reparte entre las cabezas, un número no positivo donde hace falta al menos uno, y RoPE con una dimensión de cabeza impar —la rotación va por parejas de dimensiones, así que con 63 sobra una—.

Sin esto, d_model=100, n_head=8 no falla al construir la configuración: falla mucho más tarde, dentro de un view(), con un mensaje sobre formas de tensores que no menciona ninguno de los dos números que escribiste.

src/rukh/models/config.py
@property
def head_dim(self) -> int:
"""Width of one attention head."""
return self.d_model // self.n_head
@property
def ff(self) -> int:
"""Hidden width of the MLP (``d_ff`` or ``4 * d_model``)."""
return self.d_ff if self.d_ff is not None else 4 * self.d_model

src/rukh/models/config.pylíneas 43-51 · p2

Dos propiedades derivadas y ningún campo. head_dim es 64 en small (512 / 8), 64 en tiny (256 / 4) y 64 en medium (768 / 12): los tres presets mantienen la regla de entre 64 y 128 por cabeza, por debajo de la cual una cabeza se queda sin resolución. Que sea una propiedad y no un campo es lo que impide que existan una d_model, una n_head y una head_dim que se contradigan.

src/rukh/models/config.py
PRESETS: dict[str, DecoderConfig] = {
"tiny": DecoderConfig(n_layer=6, n_head=4, d_model=256),
"small": DecoderConfig(),
"medium": DecoderConfig(n_layer=16, n_head=12, d_model=768),
}

src/rukh/models/config.pylíneas 54-58 · p2

Los tres presets del curso, en una línea cada uno, porque DecoderConfig ya trae los valores de small por defecto. Medidos: 5 309 952 parámetros y 189 segundos de entrenamiento para tiny, 38 971 392 y 42 minutos para small, 115 120 128 para medium. Los tres siguen la diagonal de las leyes de escala —profundidad y anchura creciendo juntas— en vez de estirar solo una.

src/rukh/models/config.py
def preset(name: str) -> DecoderConfig:
"""A copy of a named preset, so callers may mutate it freely."""
if name not in PRESETS:
raise ValueError(f"unknown preset {name!r}; expected one of {', '.join(PRESETS)}")
return PRESETS[name].model_copy(deep=True)

src/rukh/models/config.pylíneas 61-65 · p2

Ese model_copy(deep=True) parece paranoia y no lo es: sin él, un entrenamiento que sobreescribiera cfg.block estaría modificando el objeto del diccionario, y el siguiente preset que alguien pidiera en el mismo proceso —por ejemplo, el barrido de tres tallas de labs/m2/params.py— vendría ya contaminado. Hay un test para exactamente eso, al final de esta lección.

decoder.py: el modelo

Ciento setenta y siete líneas. Empezamos por la cabecera, que es donde se declara el contrato del fichero.

src/rukh/models/decoder.py
"""``MoveDecoder``: a hand-written GPT decoder over the fixed UCI vocabulary.
Pure model code: it knows about tensors and nothing about the data pipeline, so it can be
imported by the training loop, the sampler and the ONNX exporter alike. Blocks are pre-norm
(``x = x + attn(ln1(x))``; ``x = x + mlp(ln2(x))``), attention is causal through
``F.scaled_dot_product_attention(..., is_causal=True)``, the MLP uses GELU and the language
modelling head is tied to the token embedding.
"""
from __future__ import annotations
import math
import torch
from torch import Tensor, nn
from torch.nn import functional as F
from rukh.models.config import DecoderConfig
# ``<pad>`` is id 0 in every scheme; the loss must ignore it (see ``loader.IGNORE_INDEX``).
IGNORE_INDEX = 0

src/rukh/models/decoder.pylíneas 1-21 · p2

«Código de modelo puro: sabe de tensores y no sabe nada del pipeline de datos». Esa frase del docstring es la razón de que el bucle de entrenamiento, el muestreador y el exportador a ONNX puedan importarlo los tres sin arrastrarse unos a otros. Si MoveDecoder supiera leer un .npy, exportarlo a ONNX significaría exportar el pipeline.

IGNORE_INDEX = 0 es <pad>, y es el número más caro de olvidar de todo el módulo. La entropía cruzada del final del fichero lo recibe como ignore_index; sin eso, el modelo aprende que después de cualquier cosa viene <pad>, y lo aprende con confianza, porque es verdad de los datos que le enseñaste: todas las secuencias cortas acaban rellenas. El comentario apunta a loader.IGNORE_INDEX de M1 para que los dos no se separen nunca.

src/rukh/models/decoder.py
def rope_tables(
seq_len: int, head_dim: int, device: torch.device, base: float = 10_000.0
) -> tuple[Tensor, Tensor]:
"""``(cos, sin)`` of shape ``(seq_len, head_dim)`` for rotary position embeddings."""
inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
angles = torch.outer(torch.arange(seq_len, dtype=torch.float32), inv_freq)
full = torch.cat([angles, angles], dim=-1)
return full.cos().to(device), full.sin().to(device)

src/rukh/models/decoder.pylíneas 24-31 · p2

Las tablas de RoPE: un coseno y un seno por posición y por dimensión. inv_freq reparte 32 frecuencias entre las 64 dimensiones de una cabeza —una por pareja—, de 1 radián por posición a 0,000133. torch.outer las multiplica por cada posición, y el cat duplica el bloque de ángulos porque la variante «por mitades» que usa la función siguiente espera la tabla entera, no media.

Todo se calcula en float32 aunque el modelo entrene en 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.. Es deliberado: son ángulos acumulados a lo largo de 200 posiciones, y siete bits de mantisa no bastan para que la posición 199 tenga la fase que le toca.

src/rukh/models/decoder.py
def apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
"""Rotate ``x`` of shape ``(B, H, T, D)`` by the angles of the first ``T`` positions."""
t = x.shape[-2]
cos_t = cos[:t].to(dtype=x.dtype).view(1, 1, t, -1)
sin_t = sin[:t].to(dtype=x.dtype).view(1, 1, t, -1)
half = x.shape[-1] // 2
rotated = torch.cat([-x[..., half:], x[..., :half]], dim=-1)
return x * cos_t + rotated * sin_t

src/rukh/models/decoder.pylíneas 34-41 · p2

La rotación, escrita para las 32 parejas a la vez. cos[:t] recorta la tabla a la longitud real de la secuencia, que puede ser menor que block —un prompt de diez jugadas no paga por doscientas posiciones—, y el .to(dtype=x.dtype) baja los ángulos a la precisión del tensor después de haberlos calculado bien.

src/rukh/models/decoder.py
class CausalSelfAttention(nn.Module):
"""Multi-head causal self-attention with a single fused ``qkv`` projection."""
def __init__(self, cfg: DecoderConfig) -> None:
super().__init__()
self.n_head = cfg.n_head
self.head_dim = cfg.head_dim
self.dropout = cfg.dropout
self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model)
self.proj = nn.Linear(cfg.d_model, cfg.d_model)
self.resid_drop = nn.Dropout(cfg.dropout)

src/rukh/models/decoder.pylíneas 44-54 · p2

Una sola matriz qkv de d_model × 3·d_model produce consulta, clave y valor de golpe. Las cabezas no son parámetros: son una reinterpretación de las columnas de esa matriz, y por eso ocho cabezas cuestan exactamente lo que cuesta una. resid_drop es el dropout que se aplica a lo que la atención devuelve a la corriente residual; con dropout: 0.0 es la identidad, y está ahí porque el encoder de M3 sí lo enciende.

src/rukh/models/decoder.py
def forward(self, x: Tensor, cos: Tensor | None = None, sin: Tensor | None = None) -> Tensor:
batch, seq, _ = x.shape
q, k, v = self.qkv(x).split(x.shape[-1], dim=2)
shape = (batch, seq, self.n_head, self.head_dim)
q = q.view(shape).transpose(1, 2)
k = k.view(shape).transpose(1, 2)
v = v.view(shape).transpose(1, 2)
if cos is not None and sin is not None:
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
out = F.scaled_dot_product_attention(
q, k, v, dropout_p=self.dropout if self.training else 0.0, is_causal=True
)
out = out.transpose(1, 2).contiguous().view(batch, seq, -1)
return self.resid_drop(self.proj(out))

src/rukh/models/decoder.pylíneas 56-70 · p2

El forward de la atención es un baile de formas y merece leerse despacio, porque es donde se equivoca todo el mundo la primera vez.

self.qkv(x) da (B, T, 3d) y .split(d, dim=2) lo corta en tres de (B, T, d). El view parte la última dimensión en (n_head, head_dim) y el transpose(1, 2) mueve las cabezas delante del tiempo: (B, H, T, D). Ese orden no es estético. scaled_dot_product_attention multiplica las dos últimas dimensiones, así que la cabeza tiene que quedar fuera del producto para que cada una haga su propia atención; con (B, T, H, D) el producto mezclaría cabezas y posiciones y el modelo entrenaría igual de contento dando basura.

A la vuelta, transpose(1, 2) deshace el intercambio y el contiguous() es obligatorio: después de un transpose el tensor tiene la memoria en otro orden y view se niega a reinterpretar algo que no es contiguo. La alternativa, reshape, copiaría en silencio; contiguous().view() dice lo que hace.

dropout_p=self.dropout if self.training else 0.0 es la línea que evita el error clásico de aplicar dropout al evaluar: F.scaled_dot_product_attention no consulta self.training, así que hay que decírselo. Y is_causal=True es la máscara: sin argumentos que construir, sin una matriz T × T que materializar, y con la implementación fusionada disponible.

src/rukh/models/decoder.py
class Mlp(nn.Module):
"""Position-wise GELU feed-forward network."""
def __init__(self, cfg: DecoderConfig) -> None:
super().__init__()
self.fc = nn.Linear(cfg.d_model, cfg.ff)
self.proj = nn.Linear(cfg.ff, cfg.d_model)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x: Tensor) -> Tensor:
return self.drop(self.proj(F.gelu(self.fc(x))))

src/rukh/models/decoder.pylíneas 73-83 · p2

La parte aburrida y dos tercios de los parámetros de un bloque. Once líneas. F.gelu en vez de nn.GELU porque no hay estado que guardar, y la proyección de salida se llama proj a propósito: es el nombre que busca la inicialización escalada del constructor.

src/rukh/models/decoder.py
class Block(nn.Module):
"""One pre-norm transformer block."""
def __init__(self, cfg: DecoderConfig) -> None:
super().__init__()
self.ln1 = nn.LayerNorm(cfg.d_model)
self.attn = CausalSelfAttention(cfg)
self.ln2 = nn.LayerNorm(cfg.d_model)
self.mlp = Mlp(cfg)
def forward(self, x: Tensor, cos: Tensor | None = None, sin: Tensor | None = None) -> Tensor:
x = x + self.attn(self.ln1(x), cos, sin)
return x + self.mlp(self.ln2(x))

src/rukh/models/decoder.pylíneas 86-98 · p2

El bloque pre-normPre-normOrden de un bloque Transformer en el que la normalización va antes de la subcapa y su salida se suma a la entrada: `x = x + attn(ln1(x))`. Deja un camino sin obstáculos entre la pérdida y las primeras capas, y es lo que permite apilar doce bloques y entrenarlos sin trucos. La variante contraria (post-norm) es la del artículo original y necesita mucho más cuidado., en trece líneas de las que dos son el modelo entero: x + attn(ln1(x)) y x + mlp(ln2(x)). La normalización está dentro de cada rama, nunca en el camino recto, y ese camino recto es la autopista por la que el gradiente llega de la pérdida al embedding sin pasar por doce peajes.

src/rukh/models/decoder.py
class MoveDecoder(nn.Module):
"""GPT decoder over move tokens: ``forward`` returns ``(logits, loss)``.
``logits`` is always ``(B, T, vocab_size)``; ``loss`` is ``None`` unless ``targets`` is
given, in which case it is the mean cross entropy with ``ignore_index=0`` (``<pad>``).
"""
def __init__(self, cfg: DecoderConfig | None = None) -> None:
super().__init__()
self.cfg = cfg or DecoderConfig()
self.tokens = nn.Embedding(self.cfg.vocab_size, self.cfg.d_model)
self.positions = (
nn.Embedding(self.cfg.block, self.cfg.d_model) if self.cfg.pos == "learned" else None
)
self.drop = nn.Dropout(self.cfg.dropout)
self.blocks = nn.ModuleList(Block(self.cfg) for _ in range(self.cfg.n_layer))
self.ln_f = nn.LayerNorm(self.cfg.d_model)
self.lm_head = nn.Linear(self.cfg.d_model, self.cfg.vocab_size, bias=False)
if self.cfg.tie_embeddings:
self.lm_head.weight = self.tokens.weight
if self.cfg.pos == "rope":
cos, sin = rope_tables(self.cfg.block, 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(self._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/decoder.pylíneas 101-129 · p2

El constructor. De arriba abajo: la tabla de tokens, la de posiciones solo si pos == "learned" (con RoPE no hay tabla que aprender y el atributo se queda en None, que es lo que el forward consulta más abajo), el dropout de entrada, los doce bloques, el LayerNorm final, la cabeza de salida y el atado.

Tres detalles que se pagan caros si se cambian de orden:

  1. El atado va antes de self.apply(self._init_weights). lm_head.weight = self.tokens.weight no copia: hace que los dos nombres apunten al mismo tensor. Si la inicialización corriera antes, daría igual; si el atado se hiciera después de entrenar un rato, los dos tensores ya habrían divergido y atarlos tiraría uno de los dos.
  2. Las tablas de RoPE se registran como buffer persistent=False. Son derivadas de block y head_dim, así que guardarlas en el checkpoint sería guardar algo recalculable y, peor, fijar el contexto con el que se entrenó dentro del fichero de pesos.
  3. La inicialización escalada va la última. apply pone todo a Normal(0, 0.02) y después el bucle baja a 0.02 / √24 los tensores cuyo nombre acaba en proj.weight. Al revés, apply borraría el ajuste.
src/rukh/models/decoder.py
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)

src/rukh/models/decoder.pylíneas 131-138 · p2

La inicialización, que es la de GPT-2: normal de desviación 0,02 en cada Linear y cada Embedding, y sesgos a cero. Es un staticmethod porque no necesita nada del modelo y así self.apply puede pasárselo a cada submódulo. Los LayerNorm no aparecen: PyTorch ya los inicializa con ganancia 1 y sesgo 0, que es exactamente lo que se quiere.

src/rukh/models/decoder.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

src/rukh/models/decoder.pylíneas 140-145 · p2

Contar parámetros parece trivial y tiene una convención dentro. non_embedding=True —el valor por defecto— resta la tabla de posiciones, que es la regla de nanoGPT: las tablas de posiciones no participan del cómputo por token de la misma manera que el resto, y dejarlas fuera hace comparables dos modelos con contextos distintos. Fíjate en que no resta la tabla de tokens, porque con tie_embeddings esa tabla también es la capa de salida y sí hace cómputo.

src/rukh/models/decoder.py
def forward(self, idx: Tensor, targets: Tensor | None = None) -> tuple[Tensor, Tensor | None]:
"""Logits ``(B, T, V)`` for ``idx`` ``(B, T)`` and, with ``targets``, the scalar loss."""
seq = idx.shape[-1]
if seq > self.cfg.block:
raise ValueError(f"sequence of {seq} tokens is longer than block {self.cfg.block}")
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)
logits = self.lm_head(self.ln_f(x))
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.reshape(-1, logits.shape[-1]),
targets.reshape(-1),
ignore_index=IGNORE_INDEX,
)
return logits, loss

src/rukh/models/decoder.pylíneas 147-170 · p2

El forward del modelo, y la única rama de todo el fichero: posiciones aprendidas se suman al embedding antes de los bloques, RoPE no se suma a nada y viaja como (cos, sin) hasta dentro de la atención. Por eso cos = sin = None arranca a None y solo se llena en la rama de RoPE: la atención interpreta «None» como «esta configuración no usa rotación».

La guarda de la primera línea —una secuencia más larga que el bloque es un ValueError con los dos números dentro— existe porque el fallo natural sería un IndexError dentro de self.positions, que no dice nada útil. Y la pérdida se calcula aquí, dentro del modelo, en vez de fuera: es lo que permite que el bucle de entrenamiento sea _, loss = model(x, y) y que el mismo modelo sirva para entrenar y para inferir sin envoltorios.

El reshape(-1, V) aplana lote y tiempo en una sola dimensión de ejemplos, que es lo que cross_entropy espera. Y el ignore_index=IGNORE_INDEX es la línea de la que hablaba la cabecera.

src/rukh/models/decoder.py
@torch.no_grad()
def next_logits(self, idx: Tensor) -> Tensor:
"""Logits of the last step only, ``(B, vocab_size)``, cropped to the block size."""
cropped = idx[:, -self.cfg.block :]
logits, _ = self(cropped)
return logits[:, -1, :]

src/rukh/models/decoder.pylíneas 172-177 · p2

Seis líneas para la inferencia: los logits del último paso y nada más. El @torch.no_grad() es lo que impide que una partida de doscientos plies vaya acumulando un grafo de autograd que nadie va a derivar.

__init__.py: la puerta del paquete

src/rukh/models/__init__.py
"""Models written from scratch: the ``MoveDecoder`` GPT and its configuration."""
from rukh.models.config import PRESETS, DecoderConfig, preset
from rukh.models.decoder import IGNORE_INDEX, Block, CausalSelfAttention, Mlp, MoveDecoder
__all__ = [
"IGNORE_INDEX",
"PRESETS",
"Block",
"CausalSelfAttention",
"DecoderConfig",
"Mlp",
"MoveDecoder",
"preset",
]

src/rukh/models/__init__.pylíneas 1-15 · p2

Quince líneas de reexportes y ninguna decisión, salvo una que se nota en todo el resto del curso: from rukh.models import MoveDecoder funciona desde cualquier sitio, así que ningún otro fichero del proyecto importa rukh.models.decoder directamente. El día que el decoder se parta en dos ficheros —que es lo que M3 estuvo a punto de hacer—, esta lista es lo único que hay que tocar.

Los tests, que son el contrato

Quince tests, 159 líneas, y todos corren en CPU en menos de un segundo. No son una formalidad: son la definición ejecutable de qué es un MoveDecoder, y el módulo 3 los va a heredar casi enteros.

tests/unit/test_decoder.py
"""Tests for rukh.models: shapes, causality, loss masking, RoPE, presets and determinism."""
from __future__ import annotations
import pytest
import torch
from rukh.models import PRESETS, DecoderConfig, MoveDecoder, preset
pytestmark = pytest.mark.unit
TOY = DecoderConfig(vocab_size=64, n_layer=2, n_head=4, d_model=32, block=16)
def toy_model(seed: int = 0, **overrides: object) -> MoveDecoder:
"""A deterministic two-layer decoder in eval mode."""
torch.manual_seed(seed)
model = MoveDecoder(TOY.model_copy(update=overrides))
return model.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_decoder.pylíneas 1-24 · p2

Un modelo de juguete —dos capas, 32 dimensiones, vocabulario de 64— y dos ayudantes con semilla fija. Que sea de juguete es el punto: un test que tarda cuatro segundos en construir small es un test que nadie ejecuta al guardar el fichero. TOY.model_copy(update=overrides) deja que cada test cambie un campo sin repetir los otros cinco, y el .eval() apaga el dropout aunque aquí sea cero, porque un test que depende del valor por defecto de otro fichero es un test frágil.

tests/unit/test_decoder.py
def test_logits_have_batch_time_vocab_shape() -> None:
model = toy_model()
idx = toy_batch()
logits, loss = model(idx)
assert logits.shape == (2, 8, TOY.vocab_size)
assert loss is None
@pytest.mark.parametrize("pos", ["learned", "rope"])
def test_future_tokens_do_not_change_past_logits(pos: str) -> None:
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, _ = model(idx)
other, _ = model(changed)
assert torch.allclose(base[:, : cut + 1], other[:, : cut + 1], atol=1e-5)
assert not torch.allclose(base[:, cut + 1 :], other[:, cut + 1 :], atol=1e-5)

tests/unit/test_decoder.pylíneas 27-47 · p2

El primero comprueba las formas y que sin targets la pérdida es None. El segundo es el test del módulo: cambia todos los tokens posteriores a la posición 5 y exige que los logits anteriores no se muevan y que los posteriores sí. Está parametrizado sobre las dos formas de posición, porque RoPE toca la atención y podría romper la causalidad sin que el otro test se enterara.

Dos detalles de cómo está escrito. El assert not torch.equal(idx, changed) comprueba que el propio test hace algo: la fórmula (t + 7) % (V - 1) + 1 podría devolver el mismo token, y un test que compara una entrada consigo misma pasa siempre. Y el segundo assert not torch.allclose(...) sobre el futuro es el que impide la versión degenerada del test, un modelo que devuelva ceros.

tests/unit/test_decoder.py
@pytest.mark.parametrize("pos", ["learned", "rope"])
def test_rope_and_learned_share_the_output_shape(pos: str) -> None:
logits, _ = toy_model(pos=pos)(toy_batch())
assert logits.shape == (2, 8, TOY.vocab_size)
def test_targets_give_a_scalar_loss_that_ignores_padding() -> None:
model = toy_model()
idx = toy_batch()
targets = toy_batch(seed=2)
_, loss = model(idx, targets)
assert loss is not None and loss.ndim == 0 and torch.isfinite(loss)
padded = targets.clone()
padded[:, 4:] = 0 # <pad>: ignored by the loss
_, masked_loss = model(idx, padded)
_, short_loss = model(idx[:, :4], targets[:, :4])
assert masked_loss is not None and short_loss is not None
assert torch.allclose(masked_loss, short_loss, atol=1e-5)
all_pad = torch.zeros_like(targets)
_, empty = model(idx, all_pad)
assert empty is not None and torch.isnan(empty)

tests/unit/test_decoder.pylíneas 50-72 · p2

El de la pérdida enmascarada es más listo de lo que parece. Pone <pad> en la segunda mitad de los objetivos y exige que la pérdida salga idéntica a la de una secuencia recortada a la primera mitad: eso solo es cierto si ignore_index no solo ignora esas posiciones en la suma, sino también en el denominador de la media. Y la última línea fija la esquina fea: con todo relleno no hay nada que promediar y la pérdida es NaN, no cero. Que sea NaN es lo correcto —un cero sería una pérdida buenísima— y es la razón de que el bucle de entrenamiento filtre los NaN antes de registrarlos en MLflow.

tests/unit/test_decoder.py
def test_same_seed_gives_the_same_weights_and_logits() -> None:
idx = toy_batch()
with torch.no_grad():
a, _ = toy_model(seed=3)(idx)
b, _ = toy_model(seed=3)(idx)
c, _ = toy_model(seed=4)(idx)
assert torch.equal(a, b)
assert not torch.allclose(a, c, atol=1e-4)
def test_next_logits_match_the_last_step_of_forward() -> None:
model = toy_model()
idx = toy_batch(seq=9)
with torch.no_grad():
full, _ = model(idx)
assert torch.allclose(model.next_logits(idx), full[:, -1], atol=1e-6)
def test_next_logits_crop_sequences_longer_than_the_block() -> None:
model = toy_model()
long_idx = toy_batch(seq=TOY.block + 5)
assert model.next_logits(long_idx).shape == (2, TOY.vocab_size)
with pytest.raises(ValueError, match="longer than block"):
model(long_idx)

tests/unit/test_decoder.pylíneas 75-98 · p2

Determinismo con semilla, next_logits coincidiendo con el último paso del forward completo, y el recorte: una secuencia de 21 tokens con block=16 pasa por next_logits y falla por el forward directo. Los dos comportamientos son deliberados y están aquí escritos, que es la diferencia entre una decisión y una casualidad.

tests/unit/test_decoder.py
def test_embeddings_are_tied_and_can_be_untied() -> None:
tied = toy_model()
assert tied.lm_head.weight is tied.tokens.weight
untied = toy_model(tie_embeddings=False)
assert untied.lm_head.weight is not untied.tokens.weight
def test_preset_sizes(capsys: pytest.CaptureFixture[str]) -> None:
counts = {}
for name in PRESETS:
model = MoveDecoder(preset(name))
counts[name] = (model.num_params(), model.num_params(non_embedding=False))
with capsys.disabled():
print()
for name, (non_embedding, total) in counts.items():
print(f"{name:<7} {non_embedding:>12,} non-embedding {total:>12,} total")
assert 38_000_000 <= counts["small"][0] <= 45_000_000
assert counts["tiny"][0] < counts["small"][0] < counts["medium"][0]
def test_preset_copies_are_independent() -> None:
first = preset("tiny")
first.dropout = 0.5
assert preset("tiny").dropout == 0.0
assert PRESETS["tiny"].dropout == 0.0
with pytest.raises(ValueError, match="unknown preset"):
preset("enormous")

tests/unit/test_decoder.pylíneas 101-127 · p2

El atado se comprueba con is, no con ==: lo que hay que garantizar no es que los dos tensores tengan los mismos números, es que sean el mismo objeto. Con == el test pasaría también con dos copias recién inicializadas de la misma semilla, que es justo el bug que se quiere cazar.

test_preset_sizes usa capsys.disabled() para imprimir siempre las tres cuentas, incluso con pytest -q: es un test que además es un informe. Y la aserción es una horquilla generosa (38 a 45 millones) en vez de la cifra exacta, porque el número exacto ya lo comprueba labs/m2/params.py contra la fórmula cerrada y un test que se rompe al tocar el vocabulario sería ruido.

tests/unit/test_decoder.py
def test_config_rejects_unknown_keys_and_bad_shapes() -> None:
with pytest.raises(ValueError):
DecoderConfig(n_layers=3) # type: ignore[call-arg]
with pytest.raises(ValueError, match="not divisible"):
DecoderConfig(d_model=100, n_head=8)
with pytest.raises(ValueError, match="even head dimension"):
DecoderConfig(d_model=12, n_head=4, pos="rope")
def test_d_ff_defaults_to_four_times_d_model() -> None:
assert DecoderConfig(d_model=64, n_head=4).ff == 256
assert DecoderConfig(d_model=64, n_head=4, d_ff=128).ff == 128
assert MoveDecoder(TOY.model_copy(update={"d_ff": 48})).blocks[0].mlp.fc.out_features == 48

tests/unit/test_decoder.pylíneas 130-142 · p2

Las tres formas de pedir un modelo imposible, en el mismo orden que el validador, y la resolución de d_ff. Fíjate en la última línea: no comprueba la propiedad, comprueba que el modelo construido tiene un mlp.fc de 48 salidas. La diferencia entre probar la configuración y probar que la configuración se usa.

tests/unit/test_decoder.py
def test_a_few_steps_of_gradient_descent_reduce_the_loss() -> None:
torch.manual_seed(5)
model = MoveDecoder(TOY)
idx = toy_batch(seq=8, seed=6)
targets = torch.roll(idx, -1, dims=1)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-2)
losses = []
for _ in range(10):
_, loss = model(idx, targets)
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_decoder.pylíneas 145-159 · p2

Y el último, que es el que da más tranquilidad por línea escrita: diez pasos de AdamW sobre un lote de juguete tienen que bajar la pérdida. No comprueba que el modelo sea bueno; comprueba que los gradientes llegan a todas partes. Un tensor desconectado del grafo —un .detach() de más, un torch.no_grad() mal puesto, un with que sobra— no falla ningún otro test de este fichero y hace que el modelo no aprenda nunca.

// Ejercicio 01Rompe el atado y mira qué test se entera

En una copia del repositorio, cambia self.lm_head.weight = self.tokens.weight por self.lm_head.weight = self.tokens.weight.clone() y ejecuta uv run pytest -m unit -q. ¿Qué test falla y cuál no? Después ejecuta uv run python labs/m2/params.py: ¿cuántos parámetros tiene ahora small y por qué el assert del script no salta?

// SoluciónVer la solución

Falla test_embeddings_are_tied_and_can_be_untied, en el assert tied.lm_head.weight is tied.tokens.weight, y no falla ningún otro: el modelo construye, entrena, muestrea y exporta exactamente igual. Es el mejor ejemplo del módulo de por qué un test de identidad vale la pena: el bug no produce un error, produce otro modelo, con un millón de parámetros más que se entrenan por separado.

params.py sí cambia de número —2 030 × 512 = 1 039 360 parámetros más, 40 010 752 en total— y el assert no salta, porque la fórmula del script consulta cfg.tie_embeddings, que sigue siendo True… y el modelo ya no lo está. Es la contrapartida honesta del ejercicio: una fórmula que lee la configuración comprueba que la configuración y el modelo coinciden solo mientras el modelo obedezca la configuración. La identidad de los tensores hay que comprobarla mirando los tensores.

Qué has aprendido

src/rukh/models/ entero: una configuración que se niega a describir un modelo imposible, 177 líneas de decoder con una atención fusionada, una inicialización que tiene en cuenta cuántas veces se escribe en la corriente residual, y quince tests que son su contrato.

Cómo se mide: uv run pytest -m unit -q tests/unit/test_decoder.py pasa los quince, y uv run python labs/m2/params.py —el lab de la lección 10— imprime 38 971 392 para small y comprueba que la fórmula escrita a mano cuadra con el modelo construido. Si las dos cifras no coinciden, tienes una idea equivocada de tu propio modelo.

Lo siguiente es hacerlo converger: el bucle, el planificador de la tasa de aprendizaje, los checkpoints con procedencia y las dos configuraciones de entrenamiento del repositorio.