// 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».
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.
"""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 BaseConfigBaseConfig 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ó.
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 = TrueTrece líneas y ocho decisiones. Las que no son evidentes:
d_ff: int | None = Noneen vez de2048. El ancho del MLP casi siempre es4 * d_model, y escribirlo como un número obligaría a cambiar dos campos para cambiar la talla del modelo.Nonesignifica «el que toca», y la propiedadffde 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 = 200es 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.
@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 selfUn 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.
@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_modelDos 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.
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),}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.
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)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.
"""``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 beimported 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 languagemodelling head is tied to the token embedding."""
from __future__ import annotations
import math
import torchfrom torch import Tensor, nnfrom 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«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.
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)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.
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_tLa 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.
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)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.
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))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.
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))))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.
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))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.
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)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:
- El atado va antes de
self.apply(self._init_weights).lm_head.weight = self.tokens.weightno 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. - Las tablas de RoPE se registran como buffer
persistent=False. Son derivadas deblockyhead_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. - La inicialización escalada va la última.
applypone todo aNormal(0, 0.02)y después el bucle baja a0.02 / √24los tensores cuyo nombre acaba enproj.weight. Al revés,applyborraría el ajuste.
@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)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.
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 totalContar 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.
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, lossEl 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.
@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, :]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
"""Models written from scratch: the ``MoveDecoder`` GPT and its configuration."""
from rukh.models.config import PRESETS, DecoderConfig, presetfrom rukh.models.decoder import IGNORE_INDEX, Block, CausalSelfAttention, Mlp, MoveDecoder
__all__ = [ "IGNORE_INDEX", "PRESETS", "Block", "CausalSelfAttention", "DecoderConfig", "Mlp", "MoveDecoder", "preset",]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 for rukh.models: shapes, causality, loss masking, RoPE, presets and determinism."""
from __future__ import annotations
import pytestimport 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)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.
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)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.
@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)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.
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)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.
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")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.
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 == 48Las 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.
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]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.