// 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`.
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.
"""Transformer building blocks shared by the causal decoder and the bidirectional encoder.
The only structural difference between ``MoveDecoder`` and ``PositionEncoder`` is the attentionmask: the decoder may look left, the encoder may look everywhere. Everything else (pre-normresidual blocks, the fused ``qkv`` projection, the GELU MLP, rotary embeddings) is literally thesame code, so it lives here once and both models import it. The classes take plain numbersrather 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 checkpointwritten so far, and changing them would silently break ``load_state``."""
from __future__ import annotations
import torchfrom torch import Tensor, nnfrom torch.nn import functional as FLos 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.
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_tEs 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.
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)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.
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))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
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))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.
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)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.
from rukh.models import layersfrom rukh.models.config import DecoderConfigfrom rukh.models.layers import apply_rope, rope_tables
__all__ = [ "IGNORE_INDEX", "Block", "CausalSelfAttention", "Mlp", "MoveDecoder", "apply_rope", "rope_tables",]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.
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)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.
@staticmethod def _init_weights(module: nn.Module) -> None: layers.init_weights(module)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 torchfrom 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:
from rukh.config import BaseConfigfrom rukh.models.squares import SQUARE_TOKENS, SQUARE_VOCAB_SIZEconfig.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.
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``."""Los campos son los quince millones de parámetros de la lección 1 escritos como datos. Tres merecen comentario:
inputes unLiteralde dos valores. Un YAML que digainput: squarefalla al cargar con un mensaje que lista los válidos, y comoEncoderConfighereda deBaseConfig(conextra="forbid", la decisión de M0) una clave mal escrita tampoco pasa. Fallar al cargar, no al usar.- Dos tamaños de vocabulario,
vocab_sizeysquare_vocab, en vez de uno que signifique lo que toque: el que se usa lo eligeinput, y unvocab_size: 2030heredado de una plantilla no confunde a nadie. Lo mismo conblock, que solo pinta enmoves. 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.
@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 selfCada 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.
@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_TOKENStokens 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
"""Models written from scratch: the ``MoveDecoder`` GPT, the ``PositionEncoder`` and configs."""
from rukh.models.config import PRESETS, DecoderConfig, EncoderConfig, presetfrom rukh.models.decoder import IGNORE_INDEX, Block, CausalSelfAttention, Mlp, MoveDecoderfrom rukh.models.encoder import MMM_IGNORE_INDEX, PositionEncoderfrom 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",]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.