rukh · lab

// M3 · lección 02

Las capas compartidas: un bloque, dos modelos

El refactor con el que empieza M3: sacar de `decoder.py` la atención, el MLP y el bloque pre-norm a `models/layers.py` sin cambiar una sola clave de los checkpoints de M2, y añadir `EncoderConfig` al lado de `DecoderConfig`.

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

Lección 2 de 11 del módulo «El encoder». Viene de «La línea que lo cambia todo» y sigue en «El encoder a mano».

Qué vas a construir

Un fichero nuevo, src/rukh/models/layers.py, con la atención, el MLP y el bloque pre-norm que el decoderDecoderArquitectura Transformer con atención causal: cada posición solo mira hacia atrás, así que sirve para generar de izquierda a derecha. GPT es un decoder; MoveDecoder, el modelo de Rukh, también. de M2 tenía dentro, más un EncoderConfig al lado del DecoderConfig. No hay conceptos nuevos: es la fontanería que hace que el encoderEncoderTransformer con atención bidireccional: cada posición ve toda la secuencia. No genera; representa. En Rukh el encoder lee una partida entera y produce un vector por posición del que salen el valor y la detección de errores. de la lección siguiente sea de verdad el mismo bloque con otro booleano.

Merece una lección propia porque este refactor podía romper en silencio todos los checkpoints de M2. Lo que se aprende aquí vale para cualquier refactor de un modelo que ya tiene pesos entrenados.

layers.py: el bloque, sin saber de quién es

El fichero abre con el contrato.

src/rukh/models/layers.py
"""Transformer building blocks shared by the causal decoder and the bidirectional encoder.
The only structural difference between ``MoveDecoder`` and ``PositionEncoder`` is the attention
mask: the decoder may look left, the encoder may look everywhere. Everything else (pre-norm
residual blocks, the fused ``qkv`` projection, the GELU MLP, rotary embeddings) is literally the
same code, so it lives here once and both models import it. The classes take plain numbers
rather than a config object precisely so neither model has to depend on the other's config.
Module attribute names (``ln1``, ``attn.qkv``, ``attn.proj``, ``ln2``, ``mlp.fc``, ``mlp.proj``)
and their creation order are part of the on-disk format: they are the keys of every checkpoint
written so far, and changing them would silently break ``load_state``.
"""
from __future__ import annotations
import torch
from torch import Tensor, nn
from torch.nn import functional as F

src/rukh/models/layers.pylíneas 1-18 · p3

Los bloques reciben números sueltos y no un objeto de configuración. Si Block recibiera un DecoderConfig, el encoder tendría que construirse uno falso para llamarlo. Con cinco números —d_model, n_head, ff, dropout, causal— las dos configuraciones pueden divergir cuanto quieran.

El segundo párrafo es el aviso. Un checkpointCheckpointFotografía de un entrenamiento guardada en disco: pesos, estado del optimizador, paso alcanzado, configuración y procedencia (hash del vocabulario, del manifiesto de datos y SHA de git). Sirve para reanudar, para evaluar y para publicar; en Rukh se escribe uno cada 1 000 pasos más el mejor por pérdida de validación. de PyTorch es un diccionario cuyas claves son los nombres de los atributos por los que se llega a cada tensor: blocks.3.attn.qkv.weight es una dirección. Renombrar qkv a in_proj no da ningún error al arrancar: load_state_dict con strict=False carga lo que reconoce y deja el resto con pesos aleatorios, y sale un modelo que produce números creíbles y ha olvidado media red. Por eso este refactor mueve el código sin tocar ni un nombre ni el orden de creación.

Las posiciones rotatorias van primeras porque no dependen de nada.

src/rukh/models/layers.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)
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/layers.pylíneas 21-38 · p3

Es el código de 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). de M2, movido sin tocar un carácter. Las tablas se calculan en float32 y se convierten al tipo de x en cada llamada, porque en formato 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. el ángulo se redondearía en la tabla y el error se acumularía a lo largo de la secuencia.

La atención, con la máscara como parámetro

Aquí está el cambio real del fichero respecto al de M2: dos argumentos nuevos, causal en el constructor y attn_mask en el forward.

src/rukh/models/layers.py
class SelfAttention(nn.Module):
"""Multi-head self-attention with a single fused ``qkv`` projection.
With ``causal=True`` the attention is the decoder's: ``is_causal=True`` and no mask tensor,
which is the fast SDPA path. With ``causal=False`` it is the encoder's: every position sees
every other one, except the keys an ``attn_mask`` switches off.
"""
def __init__(self, d_model: int, n_head: int, dropout: float = 0.0, causal: bool = True):
super().__init__()
self.n_head = n_head
self.head_dim = d_model // n_head
self.dropout = dropout
self.causal = causal
self.qkv = nn.Linear(d_model, 3 * d_model)
self.proj = nn.Linear(d_model, d_model)
self.resid_drop = nn.Dropout(dropout)

src/rukh/models/layers.pylíneas 41-57 · p3

causal=True es el valor por defecto a propósito: el código que ya existía sigue comportándose igual sin pasar nada. Cambiar el valor por defecto de un booleano que decide si un modelo ve el futuro es un cambio que no se puede revisar leyendo el diff.

src/rukh/models/layers.py
def forward(
self,
x: Tensor,
cos: Tensor | None = None,
sin: Tensor | None = None,
attn_mask: 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,
attn_mask=attn_mask,
dropout_p=self.dropout if self.training else 0.0,
is_causal=self.causal and attn_mask is None,
)
out = out.transpose(1, 2).contiguous().view(batch, seq, -1)
return self.resid_drop(self.proj(out))

src/rukh/models/layers.pylíneas 59-84 · p3

La línea importante es is_causal=self.causal and attn_mask is None, y el and lo exige la API: F.scaled_dot_product_attention no admite is_causal=True junto con un attn_mask. Con los dos, el comportamiento depende del backend que se elija en tiempo de ejecución: en el mejor caso lanza, en el peor aplica solo uno. Con el and, si hay máscara, la máscara manda. El decoder nunca pasa una y se queda en el camino rápido; el encoder pasa la del relleno y is_causal se apaga solo.

Pedir la atención causal con el booleano, en vez de pasar una matriz triangular como attn_mask, también es cuestión de coste. Con el booleano, el kernel fusionado se salta los bloques del triángulo superior; con la matriz, calcula la puntuación entera y luego le suma -inf. Es la mitad del trabajo frente al trabajo entero, así que el decoder no se pasa a la rama de máscara «para unificar».

El dropout se apaga fuera de train (self.training), y en la lección 6 eso importa: un encoder congelado en modo train daría un vector distinto para la misma posición en cada época.

El MLP y el bloque, que no cambian nada

src/rukh/models/layers.py
class Mlp(nn.Module):
"""Position-wise GELU feed-forward network."""
def __init__(self, d_model: int, ff: int, dropout: float = 0.0) -> None:
super().__init__()
self.fc = nn.Linear(d_model, ff)
self.proj = nn.Linear(ff, d_model)
self.drop = nn.Dropout(dropout)
def forward(self, x: Tensor) -> Tensor:
return self.drop(self.proj(F.gelu(self.fc(x))))
class Block(nn.Module):
"""One pre-norm transformer block: ``x + attn(ln1(x))`` then ``x + mlp(ln2(x))``."""
def __init__(
self, d_model: int, n_head: int, ff: int, dropout: float = 0.0, causal: bool = True
) -> None:
super().__init__()
self.ln1 = nn.LayerNorm(d_model)
self.attn = SelfAttention(d_model, n_head, dropout, causal=causal)
self.ln2 = nn.LayerNorm(d_model)
self.mlp = Mlp(d_model, ff, dropout)
def forward(
self,
x: Tensor,
cos: Tensor | None = None,
sin: Tensor | None = None,
attn_mask: Tensor | None = None,
) -> Tensor:
x = x + self.attn(self.ln1(x), cos, sin, attn_mask)
return x + self.mlp(self.ln2(x))

src/rukh/models/layers.pylíneas 87-120 · p3

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. de M2, palabra por palabra, con los submódulos creados en el mismo orden —ln1, attn, ln2, mlp— para que las claves de los checkpoints no cambien. attn_mask se añade al final de la firma y con valor por defecto None: un argumento posicional nuevo en medio convertiría una llamada antigua que pasaba cos y sin por posición en una que pasa cos como máscara, y eso tampoco da error, da basura.

src/rukh/models/layers.py
def init_weights(module: nn.Module) -> None:
"""Normal(0, 0.02) on every ``Linear`` and ``Embedding``; biases start at zero."""
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/layers.pylíneas 123-130 · p3

La inicialización de GPT-2, ahora compartida: std = 0.02, sesgos a cero. Si los dos modelos se inicializaran distinto, la diferencia aparecería en las métricas y se atribuiría a la máscara.

decoder.py, después del refactor: tres clases que traducen su configuración a layers

El fichero de M2 pierde 68 líneas y gana 23: lo que antes era el cuerpo de la atención, el MLP y el bloque ahora es una llamada al constructor de la clase de layers que le corresponde.

src/rukh/models/decoder.py
from rukh.models import layers
from rukh.models.config import DecoderConfig
from rukh.models.layers import apply_rope, rope_tables
__all__ = [
"IGNORE_INDEX",
"Block",
"CausalSelfAttention",
"Mlp",
"MoveDecoder",
"apply_rope",
"rope_tables",
]

src/rukh/models/decoder.pylíneas 22-34 · p3

Ese __all__ cuida la compatibilidad con el código. En M2 hay tests y un exportador que hacen from rukh.models.decoder import rope_tables; reexportarlos desde donde estaban evita tocar esos ficheros, y ponerlos en __all__ impide que ruff los borre como imports sin usar.

src/rukh/models/decoder.py
class CausalSelfAttention(layers.SelfAttention):
"""Multi-head causal self-attention with a single fused ``qkv`` projection."""
def __init__(self, cfg: DecoderConfig) -> None:
super().__init__(cfg.d_model, cfg.n_head, cfg.dropout, causal=True)
class Mlp(layers.Mlp):
"""Position-wise GELU feed-forward network."""
def __init__(self, cfg: DecoderConfig) -> None:
super().__init__(cfg.d_model, cfg.ff, cfg.dropout)
class Block(layers.Block):
"""One pre-norm transformer block."""
def __init__(self, cfg: DecoderConfig) -> None:
super().__init__(cfg.d_model, cfg.n_head, cfg.ff, cfg.dropout, causal=True)

src/rukh/models/decoder.pylíneas 40-58 · p3

Tres clases que solo traducen un DecoderConfig en los números que espera layers. MoveDecoder podría construir layers.Block directamente, como hace el encoder, pero se quedan por compatibilidad: from rukh.models import Block, CausalSelfAttention, Mlp es una importación que M2 usó en sus tests y en su lección, y borrarla habría convertido un refactor interno en un cambio de API pública.

La herencia en vez de la composición también es deliberada. Un self.inner = layers.SelfAttention(…) habría metido un nivel más en las claves del checkpoint (attn.inner.qkv.weight) y habría roto lo que el docstring del fichero nuevo avisa.

src/rukh/models/decoder.py
@staticmethod
def _init_weights(module: nn.Module) -> None:
layers.init_weights(module)

src/rukh/models/decoder.pylíneas 91-93 · p3

El método se queda como delegación de una línea porque _init_weights aparece en un test de M2, y porque es el sitio donde el decoder tendría una inicialización propia si algún día la quiere.

// Ejercicio 01Comprueba que el refactor no movió ninguna clave

Con el repositorio en p3, carga un checkpoint de M2 y compara sus claves con las del modelo construido por el código nuevo. <fecha> es la marca de tiempo de la carpeta que escribió tu entrenamiento de small (ls -td checkpoints/small-*/ | head -1 te da la última):

import torch
from rukh.models import DecoderConfig, MoveDecoder, preset
payload = torch.load("checkpoints/small-<fecha>/best.pt", map_location="cpu", weights_only=False)
saved = set(payload["model_state"])
built = set(MoveDecoder(preset("small")).state_dict())
print(sorted(saved - built)[:5], sorted(built - saved)[:5])

¿Qué deberían imprimir las dos listas, y qué habría pasado si layers.Block hubiera llamado a su LayerNorm norm1 en vez de ln1?

// SoluciónVer la solución

Las dos listas tienen que salir vacías: el refactor mueve código, no nombres. Si ln1 se hubiera llamado norm1, saved - built tendría blocks.0.ln1.weight, blocks.0.ln1.bias y sus veintidós hermanos (dos por cada uno de los doce bloques de small), y built - saved los mismos con el nombre nuevo.

Lo peligroso es lo que pasa si nadie ejecuta esta comprobación. load_state_dict(..., strict=True) lanza, y eso sería un buen día. Con strict=False —lo que hace falta para cargar un checkpoint al que se le han añadido cabezas, como en M3— las claves que faltan se quedan con su inicialización aleatoria. El modelo carga, entrena y da números. Solo son los números de otro modelo.

EncoderConfig: la configuración del hermano

config.py gana un import y una clase. El import decide una dependencia entre módulos:

src/rukh/models/config.py
from rukh.config import BaseConfig
from rukh.models.squares import SQUARE_TOKENS, SQUARE_VOCAB_SIZE

src/rukh/models/config.pylíneas 9-10 · p3

config.py importa de squares.py y no al revés, así que el tamaño del vocabulario de casillas es un valor calculado en vez de un 47 escrito a mano que se quedaría desfasado el día que el esquema gane un token.

src/rukh/models/config.py
class EncoderConfig(BaseConfig):
"""Shape of a ``PositionEncoder``.
``input`` picks the representation the encoder reads: ``moves`` reuses the P1 UCI
vocabulary and the ``block`` context of the decoder, ``squares`` reads the 69 fixed tokens
of ``rukh.models.squares``. The rest mirrors ``DecoderConfig``, on purpose: the two models
are the same blocks with a different attention mask.
"""
input: Literal["moves", "squares"] = "moves"
vocab_size: int = 2030 # "moves": the P1 vocabulary
square_vocab: int = SQUARE_VOCAB_SIZE # "squares": <pad>/<mask>/<cls>, pieces, turn, ...
n_layer: int = 8
n_head: int = 6
d_model: int = 384
d_ff: int | None = None # None -> 4 * d_model
block: int = 200 # "moves"; "squares" always uses its 69 fixed positions
dropout: float = 0.1
pos: Literal["learned", "rope"] = "learned"
tie_embeddings: bool = True
"""Tie the masked-move head to the token embedding; only ever applied to ``moves``."""

src/rukh/models/config.pylíneas 69-89 · p3

Los campos son los quince millones de parámetros de la lección 1 escritos como datos. Tres merecen comentario:

  • input es un Literal de dos valores. Un YAML que diga input: square falla al cargar con un mensaje que lista los válidos, y como EncoderConfig hereda de BaseConfig (con extra="forbid", la decisión de M0) una clave mal escrita tampoco pasa. Fallar al cargar, no al usar.
  • Dos tamaños de vocabulario, vocab_size y square_vocab, en vez de uno que signifique lo que toque: el que se usa lo elige input, y un vocab_size: 2030 heredado de una plantilla no confunde a nadie. Lo mismo con block, que solo pinta en moves.
  • dropout: 0.1, mientras el decoder entrena con cero: el afinado verá 438 093 posiciones, dos órdenes de magnitud menos que los tokens del preentrenamiento, y con tan pocos datos un poco de regularización paga.
src/rukh/models/config.py
@model_validator(mode="after")
def _check(self) -> EncoderConfig:
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.square_vocab, self.n_layer, self.n_head, self.d_model) < 1:
raise ValueError("vocab_size, square_vocab, n_layer, n_head and d_model must be > 0")
if self.block < 1:
raise ValueError(f"block must be positive, got {self.block}")
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 91-105 · p3

Cada comprobación corresponde a una forma de que el modelo se construya y falle más tarde en un sitio que no explica nada. Si d_model % n_head no divide, la atención revienta con un error sobre el número de elementos de un tensor, a cuatro llamadas del número mal puesto; aquí el mensaje trae los dos números. Con dropout = 1.0 el modelo entrenaría tranquilamente contra el vacío.

La de rope con dimensión impar es la que más se agradece. apply_rope parte el vector por la mitad y con una dimensión impar la división entera pierde una componente: el modelo entrena y una dimensión de cada cabeza no recibe información de posición. No falla nada; solo aprende algo peor por un motivo que nadie encontraría.

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
@property
def tokens(self) -> int:
"""Size of the vocabulary this scheme reads."""
return self.vocab_size if self.input == "moves" else self.square_vocab
@property
def seq(self) -> int:
"""Longest sequence this scheme produces: ``block`` for moves, 69 for squares."""
return self.block if self.input == "moves" else SQUARE_TOKENS

src/rukh/models/config.pylíneas 107-125 · p3

tokens y seq resuelven en un solo sitio la dualidad de los dos esquemas, así que PositionEncoder se construye con cfg.tokens y cfg.seq sin un solo if cfg.input == …. Preguntar por el esquema en cada sitio que necesita un tamaño es lo que produce el bug de construir la tabla de posiciones con 200 entradas para una secuencia de 69 y no enterarse hasta la exportación.

Los reexportes

src/rukh/models/__init__.py
"""Models written from scratch: the ``MoveDecoder`` GPT, the ``PositionEncoder`` and configs."""
from rukh.models.config import PRESETS, DecoderConfig, EncoderConfig, preset
from rukh.models.decoder import IGNORE_INDEX, Block, CausalSelfAttention, Mlp, MoveDecoder
from rukh.models.encoder import MMM_IGNORE_INDEX, PositionEncoder
from rukh.models.heads import (
HEADS,
BlunderHead,
HeadWeights,
MultiHead,
ResultHead,
ValueHead,
)
__all__ = [
"HEADS",
"IGNORE_INDEX",
"MMM_IGNORE_INDEX",
"PRESETS",
"Block",
"BlunderHead",
"CausalSelfAttention",
"DecoderConfig",
"EncoderConfig",
"HeadWeights",
"Mlp",
"MoveDecoder",
"MultiHead",
"PositionEncoder",
"ResultHead",
"ValueHead",
"preset",
]

src/rukh/models/__init__.pylíneas 1-33 · p3

Es la superficie pública del paquete, y layers no está en ella: los bloques son un detalle de implementación de los dos modelos, y quien los necesita los pide por su ruta completa.

Qué has aprendido

Compartir el bloque es lo que hace comprobable que la única diferencia entre los dos modelos sea la máscara. El precio es cuidar tres cosas que no dan error al romperse: los nombres de los atributos (son las claves del checkpoint), el orden de los argumentos nuevos (al final y opcionales) y el valor por defecto del booleano que decide si un modelo ve el futuro. Las tres valen para cualquier refactor de un modelo con pesos ya entrenados.

Cómo se mide: uv run pytest -m unit -q sigue verde con los tests de M2 sin tocar ninguno, y cargar un checkpoint de small con el código de p3 no reporta ni una clave que falte ni una sobrante.

Lo siguiente es el encoder, que con esto cabe en una clase: el mismo bucle de bloques con causal=False, la máscara del relleno, la cabeza de jugada tapada y el pooling.