Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

42. More models: a registry, scaled RoPE, sliding windows and latent attention

In this chapter

  • What actually differs between Llama, Mistral, Qwen2, Qwen3 and DeepSeek, and how one configurable decoder covers the dense ones.
  • Long-context RoPE scaling (linear, Llama 3, YaRN), derived from what each pair of features does over the training context.
  • Sliding-window attention in the engine's kernels.
  • Multi-head latent attention: DeepSeek's KV cache, 71 times smaller than full multi-head attention would be, and the "absorbed" form that serves it as one wide query head.
  • DeepSeek's grouped router with shared experts, a registry keyed on config.json, and Llama-architecture GGUF files.

You will build

rope_inv_freq, DecoderAttention.forward, MLAAttention.paged_forward and GroupedRouter.forward in engine/models.py.

Time: 5-7 hours. GPU: not needed. The parity tests need transformers (they build random models locally; nothing is downloaded).

What differs between model families

Since Chapter 17 the engine has run Qwen3. vLLM and llama.cpp each support well over a hundred architectures. That sounds like a lot of code, but most of the decoder-only models people serve differ in a short list of switches:

Llama 3Mistral 7BQwen2 / 2.5Qwen3DeepSeek-V3
Q/K/V biasesno (attention_bias adds them, and on o_proj)noQ, K, Vnono
per-head Q/K normnononoyesno (norms on the latents)
RoPELlama 3 scalingplainplain or YaRNplain or YaRNYaRN, interleaved pairs
attentionGQAGQA, 4,096-token windowGQA, optional windowsGQAmulti-head latent
MLPSwiGLUSwiGLUSwiGLUSwiGLU or MoEfirst 3 layers dense, then 256 routed + 1 shared expert
tied embeddings1B, 3Bnosmall modelssmall modelsno

Everything in the first four columns fits one class whose constructor reads these switches from config.json. DeepSeek needs a different attention module and router. Gemma (normalized embeddings, 1 + w norms, soft-capping), Phi and others need a handful more switches (stretch exercise 1).

@dataclass
class DecoderConfig(Qwen3Config):
    model_type: str = "qwen3"
    qkv_bias: bool = False            # Qwen2
    o_bias: bool = False
    qk_norm: bool = True              # Qwen3
    rope_scaling: dict | None = None
    sliding_window: int = 0
    layer_types: list | None = None   # "full_attention" / "sliding_attention" per layer

    @classmethod
    def from_hf(cls, raw):
        kind = raw.get("model_type")
        if kind not in ("llama", "mistral", "qwen2", "qwen3"):
            raise ValueError(f"DecoderConfig doesn't cover model_type {kind!r}")
        if raw.get("hidden_act", "silu") != "silu" or raw.get("mlp_bias"):
            raise ValueError("Only SwiGLU MLPs without biases are implemented")
        layers, heads = raw["num_hidden_layers"], raw["num_attention_heads"]
        window = raw.get("sliding_window") or 0
        types = raw.get("layer_types")
        if types is None and window:
            if kind == "mistral":
                types = ["sliding_attention"] * layers
            elif kind == "qwen2" and raw.get("use_sliding_window"):
                start = raw.get("max_window_layers", layers)
                types = ["full_attention" if i < start else "sliding_attention" for i in range(layers)]
        bias = bool(raw.get("attention_bias", False))
        return cls(vocab_size=raw["vocab_size"], hidden_size=raw["hidden_size"], intermediate_size=raw["intermediate_size"],
                   num_hidden_layers=layers, num_attention_heads=heads,
                   num_key_value_heads=raw.get("num_key_value_heads") or heads,
                   head_dim=raw.get("head_dim") or raw["hidden_size"] // heads, rms_norm_eps=raw.get("rms_norm_eps", 1e-6),
                   rope_theta=theta_of(raw), max_position_embeddings=raw.get("max_position_embeddings", 4096),
                   tie_word_embeddings=raw.get("tie_word_embeddings", False), model_type=kind,
                   qkv_bias=bias or kind == "qwen2", o_bias=bias, qk_norm=kind == "qwen3",
                   rope_scaling=scaling_of(raw), sliding_window=window if types and "sliding_attention" in types else 0,
                   layer_types=types)

    def window(self, layer):
        types = self.layer_types
        return self.sliding_window if types and types[layer] == "sliding_attention" else 0

Two details hide in that table. Llama’s attention_bias adds biases to all four projections, Qwen2 always has them on Q, K and V and never on o_proj; and Qwen2’s sliding windows apply only to layers from max_window_layers on. Recent transformers configs spell the result out per layer in layer_types, which the config prefers when present.

class DecoderAttention(Qwen3Attention):
    def __init__(self, cfg, window=0):
        nn.Module.__init__(self)
        self.cfg, self.sliding_window = cfg, window
        d, hq, hkv = cfg.head_dim, cfg.num_attention_heads, cfg.num_key_value_heads
        self.q_proj = nn.Linear(cfg.hidden_size, hq * d, bias=cfg.qkv_bias)
        self.k_proj = nn.Linear(cfg.hidden_size, hkv * d, bias=cfg.qkv_bias)
        self.v_proj = nn.Linear(cfg.hidden_size, hkv * d, bias=cfg.qkv_bias)
        self.o_proj = nn.Linear(hq * d, cfg.hidden_size, bias=cfg.o_bias)
        if cfg.qk_norm:                            # absent, not identity: FlatModel checks hasattr
            self.q_norm = RMSNorm(d, cfg.rms_norm_eps)
            self.k_norm = RMSNorm(d, cfg.rms_norm_eps)

    def forward(self, x, positions, rope, cache=None, layer=0, rows=None):
        """Qwen3Attention.forward with optional norms and a sliding window.  (Your engine: Chapter 42)"""
        b, t, _ = x.shape
        c = self.cfg
        q = self.q_proj(x).view(b, t, c.num_attention_heads, c.head_dim)
        k = self.k_proj(x).view(b, t, c.num_key_value_heads, c.head_dim)
        v = self.v_proj(x).view(b, t, c.num_key_value_heads, c.head_dim).transpose(1, 2)
        if hasattr(self, "q_norm"):
            q, k = self.q_norm(q), self.k_norm(k)
        cos, sin = rope
        q, k = apply_rope(q.transpose(1, 2), cos, sin), apply_rope(k.transpose(1, 2), cos, sin)
        key_positions = None
        if cache is not None:
            k, v, key_positions = cache.update(layer, k, v, positions, rows)
        allowed = None
        if self.sliding_window:
            s = k.shape[2]
            kp = key_positions if key_positions is not None else torch.arange(s, device=x.device)
            qp = positions if positions is not None else torch.arange(s - t, s, device=x.device)
            allowed = window_mask(qp, kp, self.sliding_window)
            allowed = allowed[:, None] if allowed.ndim == 3 else allowed
        y = causal_attention(q, k, v, positions, key_positions, allowed)
        return self.o_proj(y.transpose(1, 2).reshape(b, t, -1))


class DecoderLayer(Qwen3Layer):
    def __init__(self, cfg, layer, attention=None, mlp=None):
        nn.Module.__init__(self)
        self.input_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
        self.self_attn = attention or DecoderAttention(cfg, cfg.window(layer))
        self.post_attention_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
        self.mlp = mlp or SwiGLU(cfg.hidden_size, cfg.intermediate_size)


class Backbone(nn.Module):
    def __init__(self, cfg, make_layer):
        super().__init__()
        self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size)
        self.layers = nn.ModuleList(make_layer(i) for i in range(cfg.num_hidden_layers))
        self.norm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)


class Decoder(Qwen3):
    """Llama, Mistral, Qwen2 and Qwen3 in one class."""

    def __init__(self, cfg, make_layer=None):
        nn.Module.__init__(self)
        self.cfg = cfg
        self.model = Backbone(cfg, make_layer or (lambda i: DecoderLayer(cfg, i)))
        self.lm_head = nn.Linear(cfg.hidden_size, cfg.vocab_size, bias=False)
        if cfg.tie_word_embeddings:
            self.lm_head.weight = self.model.embed_tokens.weight
        self.inv_freq, self.attention_factor = rope_inv_freq(cfg.head_dim, cfg.rope_theta, cfg.rope_scaling,
                                                             cfg.max_position_embeddings)

    def rope_tables(self, positions, dtype):
        """FlatModel's hook: [N, 1, dim] tables with this model's scaling."""
        cos, sin = rope_tables(positions, self.inv_freq, self.attention_factor, dtype)
        return cos[:, None], sin[:, None]

    def forward(self, ids, cache=None, positions=None, rows=None, return_hidden=False):
        if positions is None:
            start = cache.length if cache is not None else 0
            positions = torch.arange(start, start + ids.shape[1], device=ids.device)
        x = self.model.embed_tokens(ids)
        cos, sin = rope_tables(positions, self.inv_freq, self.attention_factor, x.dtype)
        rope = (cos[None, None], sin[None, None]) if positions.ndim == 1 else (cos[:, None], sin[:, None])
        for i, layer in enumerate(self.model.layers):
            x = layer(x, positions, rope, cache, i, rows)
        hidden = self.model.norm(x)
        return hidden if return_hidden else self.lm_head(hidden)

DecoderAttention is Chapter 17’s attention with the norms optional, biases switchable and a window mask. It leaves q_norm undefined rather than an identity, because FlatModel (Chapter 31) checks hasattr(attn, "q_norm"). Decoder adds one hook, rope_tables, which FlatModel already looks for, so the serving path needs no model-specific code for any of these families.

Scaled RoPE

RoPE rotates feature pair $i$ by angle $p,\theta^{-2i/d}$ at position $p$ (Chapter 17). Pair 0 turns fast (once every $2\pi$ positions); the last pair turns slowly: for $\theta = 500{,}000$ and $d = 128$ its wavelength is about 2.6 million positions. A model trained on 8,192-token sequences has seen fast pairs turn thousands of times, but slow pairs only through a fraction of a turn. Positions past the training length push slow pairs into angles never seen in training, and quality collapses. Every long-context scheme is a way to keep angles in the trained range:

  • Linear (position interpolation) divides every frequency by the extension factor $s$: position $p$ looks like $p/s$. Simple, but fast pairs, which distinguish neighbouring tokens, lose resolution.
  • Llama 3 keeps pairs whose wavelength is shorter than (original context / high_freq_factor), divides pairs longer than (original context / low_freq_factor) by $s$, and blends linearly in between.
  • YaRN chooses the same split by counting rotations: pairs that turn more than beta_fast (32) times over the original context are kept, pairs that turn fewer than beta_slow (1) times are interpolated, with a linear ramp between. It also sharpens attention, multiplying cos and sin by $0.1 \ln s + 1$ (DeepSeek folds a version of this into the softmax scale instead).
def rope_inv_freq(dim, theta, scaling=None, max_positions=None):
    """Rotation frequency of each feature pair, with long-context scaling.  (Your engine: Chapter 42)

    Returns (inv_freq [dim/2], attention_factor), the factor multiplying cos and sin (YaRN only).
      linear  every frequency divided by `factor`: positions squeezed into the trained range
      llama3  low frequencies (wavelength > original context / low_freq_factor) divided by
              `factor`, high ones kept, a smooth blend between
      yarn    the same idea with a linear ramp over dimensions, chosen by how many full
              rotations each pair makes over the original context, plus a temperature
    """
    base = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.int64, device="cpu").float() / dim))   # cpu: even under a meta-device build
    s = scaling or {}
    kind = s.get("rope_type", s.get("type", "default"))
    if kind == "default":
        return base, 1.0
    factor = s.get("factor")
    if kind == "linear":
        return base / factor, 1.0
    if kind == "llama3":
        old, low, high = s["original_max_position_embeddings"], s["low_freq_factor"], s["high_freq_factor"]
        wavelen = 2 * math.pi / base
        scaled = torch.where(wavelen > old / low, base / factor, base)
        smooth = (old / wavelen - low) / (high - low)
        medium = (wavelen <= old / low) & (wavelen >= old / high)
        return torch.where(medium, (1 - smooth) * base / factor + smooth * base, scaled), 1.0
    if kind == "yarn":
        old = s["original_max_position_embeddings"]
        factor = factor or max_positions / old

        def mscale(scale, m=1.0):
            return 1.0 if scale <= 1 else 0.1 * m * math.log(scale) + 1.0

        attention = s.get("attention_factor")
        if attention is None:
            attention = (mscale(factor, s["mscale"]) / mscale(factor, s["mscale_all_dim"])
                         if s.get("mscale") and s.get("mscale_all_dim") else mscale(factor))

        def dim_for(rotations):                     # the pair index that turns `rotations` times over `old`
            return dim * math.log(old / (rotations * 2 * math.pi)) / (2 * math.log(theta))

        lo, hi = dim_for(s.get("beta_fast") or 32), dim_for(s.get("beta_slow") or 1)
        if s.get("truncate", True):
            lo, hi = math.floor(lo), math.ceil(hi)
        lo, hi = max(lo, 0), min(hi, dim - 1)
        ramp = ((torch.arange(dim // 2, dtype=torch.float32, device="cpu") - lo) / ((hi - lo) or 0.001)).clamp(0, 1)
        keep = 1 - ramp                             # 1: fast pairs, extrapolated; 0: slow pairs, interpolated
        return base / factor * (1 - keep) + base * keep, attention
    raise ValueError(f"RoPE type {kind!r} is not implemented")


def rope_tables(positions, inv_freq, attention_factor, dtype):
    """cos and sin [..., dim] for any positions shape."""
    angles = positions.float()[..., None] * inv_freq.to(positions.device)
    angles = torch.cat((angles, angles), dim=-1)
    return (angles.cos() * attention_factor).to(dtype), (angles.sin() * attention_factor).to(dtype)

All three change only the frequencies, so the cos/sin tables are computed once per step as before. (Qwen’s “dynamic NTK” and Phi-3’s LongRoPE are variants: the first changes $\theta$ with the sequence length, which breaks the prefix cache’s assumption that a token’s K never changes; the second learns per-pair factors.)

Sliding windows

Mistral 7B attends to the last 4,096 tokens only; Gemma 2 and 3, gpt-oss and Qwen2 alternate windowed and full layers. The engine’s kernels already take a window argument (Chapter 32): the reference backend masks keys more than $w - 1$ positions back, and the Triton kernel also starts its loop at the first block inside the oldest query’s window, so a windowed layer costs $O(w)$ per token however long the context. FlatModel.attention now passes each layer’s sliding_window through.

What the engine doesn’t do yet is free the blocks that fall out of every window. A model with only windowed layers needs just $\lceil w / \text{block size} \rceil + 1$ blocks per request; a model that mixes windowed and full layers needs both kinds. vLLM’s hybrid KV-cache manager gives each layer group its own block table so that windowed layers’ blocks can be recycled (stretch exercise 2).

Multi-head latent attention

DeepSeek-V3 has 128 attention heads with 192-wide keys and 128-wide values. Cached the ordinary way, that’s $61 \times 128 \times (192 + 128) \times 2$ bytes = 4.8 MB per token: one 32K-token conversation would need 150 GB. MLA (DeepSeek-V2, 2024) caches 68.6 KB per token instead:

  1. Each token’s hidden state is compressed to a latent $c = \mathrm{norm}(W_{dkv},x)$ of width 512.
  2. Every head’s key and value are linear functions of it: $k^{nope}_h = W^{uk}_h c$, $v_h = W^{uv}_h c$ (kv_b_proj holds all of them).
  3. RoPE can’t pass through $W^{uk}$ (rotation depends on position, the matrix doesn’t), so each token also gets one small rotary key $k^{rope}$ of width 64, shared by all heads; queries have a matching rotary part.

Only $[c, k^{rope}]$, 576 numbers, is cached. run.py models prints the arithmetic for some real configurations:

{"model": "Llama-3.1-8B", "kv_KiB_per_token": 128.0, "GiB_for_32k_tokens": 4.0, "saving_vs_full_mha": "4.0x", "sliding_window": null}
{"model": "Mistral-7B-v0.1", "kv_KiB_per_token": 128.0, "GiB_for_32k_tokens": 4.0, "saving_vs_full_mha": "4.0x", "sliding_window": 4096}
{"model": "Qwen2.5-7B", "kv_KiB_per_token": 56.0, "GiB_for_32k_tokens": 1.75, "saving_vs_full_mha": "7.0x", "sliding_window": null}
{"model": "DeepSeek-V3", "kv_KiB_per_token": 68.6, "GiB_for_32k_tokens": 2.14, "saving_vs_full_mha": "71.1x", "sliding_window": null}

A 671-billion-parameter model whose KV cache per token is about half of Llama-3.1-8B’s.

Serving it: absorbing the up-projections

Decompressing every cached token into 128 heads’ keys and values at every decode step would cost far more than reading them. The trick is associativity. A head’s score against cached token $j$ is

$$q^{nope}_h \cdot (W^{uk}_h c_j) + q^{rope}_h \cdot k^{rope}_j = \big((W^{uk}_h)^T q^{nope}_h\big) \cdot c_j + q^{rope}_h \cdot k^{rope}_j .$$

Map each query into latent space once, $\tilde q_h = (W^{uk}_h)^T q^{nope}h$, concatenate its rotary part, and attention becomes multi-query attention with a single 576-wide KV head over the cached rows. The values are the latents themselves: the output in latent space, $\sum_j a{hj} c_j$, is mapped back per head by $W^{uv}_h$ after attention. So the cache needs one tensor, whose first 512 features serve as values. That’s exactly what DeepSeek’s FlashMLA kernel consumes.

class MLAAttention(nn.Module):
    """Multi-head latent attention (DeepSeek-V2). Keys and values of all heads are linear
    functions of ONE compressed vector per token, c = norm(W_dkv x) of width kv_lora_rank, plus
    a small rotary key shared by all heads. Only [c, k_rope] is cached."""

    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        h, qk = cfg.num_attention_heads, cfg.qk_nope_head_dim + cfg.qk_rope_head_dim
        if cfg.q_lora_rank:
            self.q_a_proj = nn.Linear(cfg.hidden_size, cfg.q_lora_rank, bias=cfg.attention_bias)
            self.q_a_layernorm = RMSNorm(cfg.q_lora_rank)
            self.q_b_proj = nn.Linear(cfg.q_lora_rank, h * qk, bias=False)
        else:
            self.q_proj = nn.Linear(cfg.hidden_size, h * qk, bias=False)
        self.kv_a_proj_with_mqa = nn.Linear(cfg.hidden_size, cfg.kv_lora_rank + cfg.qk_rope_head_dim, bias=cfg.attention_bias)
        self.kv_a_layernorm = RMSNorm(cfg.kv_lora_rank)
        self.kv_b_proj = nn.Linear(cfg.kv_lora_rank, h * (cfg.qk_nope_head_dim + cfg.v_head_dim), bias=False)
        self.o_proj = nn.Linear(h * cfg.v_head_dim, cfg.hidden_size, bias=cfg.attention_bias)
        self.scale = qk ** -0.5
        s = cfg.rope_scaling or {}
        if s.get("mscale_all_dim"):                  # YaRN's temperature, applied to the softmax scale
            m = 0.1 * s["mscale_all_dim"] * math.log(s["factor"]) + 1.0 if s["factor"] > 1 else 1.0
            self.scale *= m * m

    def queries(self, x):
        c = self.cfg
        q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(x))) if c.q_lora_rank else self.q_proj(x)
        q = q.view(*x.shape[:-1], c.num_attention_heads, c.qk_nope_head_dim + c.qk_rope_head_dim)
        return q.split([c.qk_nope_head_dim, c.qk_rope_head_dim], dim=-1)

    def latent(self, x):
        c = self.cfg
        latent, k_rope = self.kv_a_proj_with_mqa(x).split([c.kv_lora_rank, c.qk_rope_head_dim], dim=-1)
        return self.kv_a_layernorm(latent), k_rope

    def forward(self, x, positions, rope, cache=None, layer=0, rows=None):
        """Training-style MLA: expand every head's K and V, then ordinary attention."""
        if cache is not None:
            raise NotImplementedError("MLA decodes through the engine's paged latent cache (Chapter 42)")
        c = self.cfg
        b, t, _ = x.shape
        q_nope, q_rope = self.queries(x)                             # [b, t, H, *]
        latent, k_rope = self.latent(x)
        cos, sin = rope
        q_rope = apply_rope(deinterleave(q_rope.transpose(1, 2)), cos, sin)
        k_rope = apply_rope(deinterleave(k_rope[:, None]), cos, sin)                 # [b, 1, t, rope]
        kv = self.kv_b_proj(latent).view(b, t, c.num_attention_heads, -1).transpose(1, 2)
        k_nope, v = kv.split([c.qk_nope_head_dim, c.v_head_dim], dim=-1)
        q = torch.cat((q_nope.transpose(1, 2), q_rope), dim=-1)
        k = torch.cat((k_nope, k_rope.expand(-1, c.num_attention_heads, -1, -1)), dim=-1)
        y = causal_attention(q, k, v, positions, scale=self.scale)
        return self.o_proj(y.transpose(1, 2).reshape(b, t, -1))

    def paged_forward(self, h, positions, rope, kv, meta, backend):
        """Serving-style MLA with the up-projections absorbed.  (Your engine: Chapter 42)

        score = q_nope . (W_uk c) + q_rope . k_rope = (W_uk^T q_nope) . c + q_rope . k_rope, so
        each query is mapped into latent space once and attends straight over the cached
        [c, k_rope] rows: multi-query attention with ONE head of width kv_lora_rank +
        qk_rope_head_dim (576 in V3). The values are the latents themselves (the first
        kv_lora_rank features of the same rows); W_uv maps the result back per head.
        """
        c = self.cfg
        n, r = h.shape[0], c.kv_lora_rank
        q_nope, q_rope = self.queries(h)                               # [N, H, *]
        latent, k_rope = self.latent(h)
        cos, sin = rope
        q_rope = apply_rope(deinterleave(q_rope), cos, sin)
        k_rope = apply_rope(deinterleave(k_rope[:, None]), cos, sin)   # [N, 1, rope]
        w = self.kv_b_proj.weight.view(c.num_attention_heads, c.qk_nope_head_dim + c.v_head_dim, r)
        q_latent = torch.einsum("nhd,hdr->nhr", q_nope, w[:, :c.qk_nope_head_dim])
        row = torch.cat((latent[:, None], k_rope), dim=-1)             # what the cache holds
        backend.write(kv, row, row, meta)
        out = backend.forward(torch.cat((q_latent, q_rope), dim=-1), kv, meta, scale=self.scale)[..., :r]
        y = torch.einsum("nhr,hvr->nhv", out, w[:, c.qk_nope_head_dim:])
        return self.o_proj(y.reshape(n, -1))

paged_forward is that computation over the paged pool: it writes the step’s $[c, k^{rope}]$ rows, runs the backend’s attention with 128 query heads against one KV head, keeps the first 512 output features and applies $W^{uv}$. forward is the training-style version (decompress, then ordinary attention), kept for parity tests. Two smaller details: DeepSeek’s checkpoints pair rotary features as adjacent elements $(0, 1), (2, 3), \ldots$, so deinterleave reorders them to the split-half pairing that apply_rope uses (applied to queries and keys alike, so dot products are unchanged). And YaRN’s temperature multiplies the softmax scale by $(0.1 \cdot \texttt{mscale_all_dim} \cdot \ln s + 1)^2$.

FlatModel.attention hands any attention module with a paged_forward method its own path. The model reports its pool shape through kv_spec (one “head” of width 576) and allocates it through allocate_kv, which returns the same tensor as both K and V, so the engine’s block manager, prefix cache, swapping and disaggregated transfer (Chapter 41) all work unchanged:

class MLADecoder(Decoder):
    """DeepSeek-V2/V3 (and Kimi K2): MLA in every layer; dense MLPs in the first
    first_k_dense_replace layers, DeepseekMoE after."""

    kv_pools = 1                       # one latent pool per layer, not separate K and V

    def __init__(self, cfg):
        def layer(i):
            moe = cfg.n_routed_experts and i >= cfg.first_k_dense_replace
            return DecoderLayer(cfg, i, MLAAttention(cfg), DeepseekMoE(cfg) if moe else None)
        super().__init__(cfg, layer)

    def kv_spec(self):
        c = self.cfg
        return c.num_hidden_layers, 1, c.kv_lora_rank + c.qk_rope_head_dim

    def allocate_kv(self, shape, dtype, device):
        """K and V are the same latent rows: one tensor serves as both."""
        pool = torch.zeros(shape, dtype=dtype, device=device)
        return pool, pool

Absorption isn’t always the cheaper form. For a long prefill, decompressing the chunk’s latents and running ordinary attention over 128 heads of width 192 does fewer FLOPs than 128 heads of width 576. vLLM and SGLang use the decompressed form for prefill and the absorbed form for decode (stretch exercise 3). MLA also doesn’t split over tensor parallelism the usual way: it has one KV head, so every rank would cache the whole latent. DeepSeek’s deployments run attention data-parallel instead (Chapter 41), and shard_tensor_parallel refuses MLA models.

DeepSeek’s router

DeepSeek-V3’s MoE layers have 256 routed experts (8 chosen per token) and one shared expert that every token uses. The router has two twists over Chapter 27’s:

  • Group-limited routing. The 256 experts are split into 8 groups. Each group is scored by the sum of its two best experts, only the top 4 groups are eligible, and the 8 experts come from those. Under expert parallelism with groups mapped to nodes, a token’s experts then span at most 4 nodes, which bounds its all-to-all traffic.
  • Auxiliary-loss-free balancing. Scores are sigmoids, not a softmax, and each expert has a bias added for selection only. During training the bias rises for underused experts and falls for overused ones; the combine weights use the raw scores, normalized over the chosen 8 and multiplied by routed_scaling_factor (2.5).

DeepSeek-V2 uses a softmax and, for its large model, groups scored by their best expert. One class covers both:

class GroupedRouter(nn.Module):
    """DeepSeek's router. Experts are split into n_group groups; only the topk_group best groups
    may be chosen (which bounds how many nodes a token's experts span under expert parallelism),
    then the top k experts among them. V3 scores with a sigmoid and adds a per-expert bias,
    adjusted during training to balance load, for SELECTION only; the weights use the raw scores."""

    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.weight = nn.Parameter(torch.zeros(cfg.n_routed_experts, cfg.hidden_size))
        if cfg.scoring == "sigmoid":
            self.e_score_correction_bias = nn.Parameter(torch.zeros(cfg.n_routed_experts), requires_grad=False)

    def forward(self, x):
        """x [N, D] -> (logits, weights [N, k], experts [N, k]).  (Your engine: Chapter 42)"""
        c = self.cfg
        logits = F.linear(x.float(), self.weight.float())
        scores = logits.sigmoid() if c.scoring == "sigmoid" else logits.softmax(-1)
        choice = scores + self.e_score_correction_bias if c.scoring == "sigmoid" else scores
        if c.n_group > 1:
            grouped = choice.view(-1, c.n_group, c.n_routed_experts // c.n_group)
            group_score = grouped.topk(2, dim=-1).values.sum(-1) if c.group_score == "top2" else grouped.amax(-1)
            keep = torch.zeros_like(group_score).scatter_(1, group_score.topk(c.topk_group, dim=-1).indices, 1)
            choice = choice.masked_fill(~keep.bool().repeat_interleave(c.n_routed_experts // c.n_group, dim=1), float("-inf"))
        experts = choice.topk(c.num_experts_per_tok, dim=-1).indices
        weights = scores.gather(1, experts)
        if c.norm_topk_prob:
            weights = weights / (weights.sum(-1, keepdim=True) + 1e-20)
        return logits, (weights * c.routed_scaling_factor).to(x.dtype), experts


class DeepseekMoE(nn.Module):
    """Routed experts (Chapter 27's stacked layout) plus always-on shared experts."""

    def __init__(self, cfg):
        super().__init__()
        self.gate = GroupedRouter(cfg)
        self.experts = Experts(cfg.n_routed_experts, cfg.hidden_size, cfg.moe_intermediate_size)
        self.shared_experts = (SwiGLU(cfg.hidden_size, cfg.moe_intermediate_size * cfg.n_shared_experts)
                               if cfg.n_shared_experts else None)

    def forward(self, x):
        shape = x.shape
        flat = x.reshape(-1, shape[-1])
        _, weights, experts = self.gate(flat)
        out = self.experts.forward_grouped(flat, weights, experts)
        if self.shared_experts is not None:
            out = out + self.shared_experts(flat)
        return out.reshape(shape)

The registry

def build_model(raw):
    """An empty model for a config.json dict."""
    kind = raw.get("model_type")
    if kind in ("llama", "mistral", "qwen2", "qwen3"):
        return Decoder(DecoderConfig.from_hf(raw))
    if kind == "qwen3_moe":
        from .moe import Qwen3Moe, Qwen3MoeConfig
        return Qwen3Moe(Qwen3MoeConfig.from_hf(raw))
    if kind in ("deepseek_v2", "deepseek_v3", "kimi_k2"):
        return MLADecoder(MLAConfig.from_hf({**raw, "model_type": "deepseek_v3" if kind == "kimi_k2" else kind}))
    raise ValueError(f"No model registered for model_type {kind!r}")


@torch.no_grad()
def load_weights(model, named_tensors):
    """Name-to-name copy. Per-expert checkpoint tensors (experts.E.gate_proj.weight) are stacked
    into the [E, 2I, D] / [E, D, I] layout; layers past num_hidden_layers (DeepSeek-V3's
    multi-token-prediction module) are skipped."""
    from .loaders import assign
    from .moe import stack_expert_tensors
    cfg = model.cfg
    params = dict(model.named_parameters())
    loaded, pending = set(), {}
    for name, value in named_tensors:
        parts = name.split(".")
        if name.startswith("model.layers.") and int(parts[2]) >= cfg.num_hidden_layers:
            continue
        if ".mlp.experts." in name and parts[5].isdigit():
            pending.setdefault(int(parts[2]), {})[(int(parts[5]), parts[6])] = value
            continue
        if name == "lm_head.weight" and cfg.tie_word_embeddings:
            continue
        if name not in params:
            raise ValueError(f"Unexpected tensor {name}")
        assign(params[name], value, name)
        loaded.add(name)
    for layer, tensors in pending.items():
        experts = model.model.layers[layer].mlp.experts
        e, inter = experts.gate_up_proj.shape[0], experts.down_proj.shape[2]
        gate_up, down = stack_expert_tensors(tensors, e, inter)
        assign(experts.gate_up_proj, gate_up, f"layer {layer} experts")
        assign(experts.down_proj, down, f"layer {layer} experts")
        loaded |= {f"model.layers.{layer}.mlp.experts.gate_up_proj", f"model.layers.{layer}.mlp.experts.down_proj"}
    missing = set(params) - loaded - ({"lm_head.weight"} if cfg.tie_word_embeddings else set())
    if missing:
        raise ValueError(f"Missing tensors: {sorted(missing)[:5]}")
    return model.eval()


def load_registered(directory, device="cpu", dtype=torch.bfloat16):
    """Any registered architecture from a local snapshot (Chapter 18's loader, generalized)."""
    from .loaders import read_config
    from .safetensors_io import snapshot_tensors
    raw = read_config(directory)
    if raw.get("quantization_config"):
        raise ValueError("Quantized checkpoints: see formats.hf_quant (Chapter 39)")
    with torch.device("meta"):
        model = build_model(raw)
    model = model.to_empty(device=device).to(dtype)
    if model.cfg.tie_word_embeddings:
        model.lm_head.weight = model.model.embed_tokens.weight
    return load_weights(model, snapshot_tensors(directory))

build_model maps config.json’s model_type to a class; load_weights is Chapter 18’s name-to-name copy, plus stacking for per-expert tensors (experts.17.gate_proj.weight) and skipping DeepSeek-V3’s multi-token-prediction layer, which follows the last regular layer (Chapter 37 could use it as a drafter). loaders.load_model now falls through to this registry for anything it doesn’t handle itself. Kimi K2 declares model_type: kimi_k2 with DeepSeek-V3’s architecture, so it maps to the same class.

GGUF files get the same coverage. Llama-architecture files (which include Mistral) store Q and K rows permuted for llama.cpp’s interleaved RoPE, so load_gguf undoes the permutation (Chapter 38), on quantized rows too: GGML quantizes each row independently, so permuting rows of codes and scales is exact. Llama 3’s frequency scaling arrives as a rope_freqs tensor of per-pair divisors; Qwen2 files carry attn_q.bias tensors. export_gguf writes both architectures, which the round-trip test uses.

Build it

Engine milestone 42: more models. Implement rope_inv_freq, DecoderAttention.forward, MLAAttention.paged_forward and GroupedRouter.forward in engine/models.py.

pytest tests/test_ch42_models.py
python run.py models

The tests compare logits with transformers for Llama 3 (scaled RoPE, tied embeddings), Llama with linear scaling and biases, Mistral (sliding window), Qwen2 (Q/K/V biases, windows on later layers, YaRN), DeepSeek-V3 (MLA with a low-rank query, YaRN, grouped sigmoid routing) and DeepSeek-V2-Lite (plain query, softmax routing, two shared experts); check that the engine generates the same tokens as a cache-free forward for each; load checkpoints written by save_pretrained; check the latent pool’s size; and round-trip Llama and Qwen2 GGUF files.

What hasn’t been checked here: real checkpoints. The test environment can’t download from the Hugging Face Hub, so every comparison uses randomly initialized models with the real architectures (Appendix F). Load a real Llama-3.2-1B or DeepSeek-V2-Lite and compare greedy outputs with transformers before trusting a new family.

Stretch exercises

  1. ★★ Add Gemma 3: RMSNorm with (1 + weight), embeddings scaled by $\sqrt{d}$, GeGLU MLPs, Q/K norms, an extra norm after each sublayer, and five windowed layers for every full one. Compare with transformers. Where: configuration, layers and registry in engine/models.py, with execution support in engine/serve/model.py.
  2. ★★★ Free out-of-window blocks: give windowed and full layers separate block tables in BlockManager, and recycle a request’s windowed-layer blocks once every query has moved past them. How many more Mistral requests fit in the same pool at 32K context? Where: BlockManager in engine/serve/blocks.py, with per-layer tables in engine/serve/batch.py and attention dispatch in engine/serve/model.py.
  3. ★★ Implement MLA’s prefill path: decompress the context’s latents for a long prompt and run ordinary attention with 128 heads. At what prompt length does it beat the absorbed form on your hardware? Where: MLAAttention.forward / paged_forward in engine/models.py.
  4. ★★ Load DeepSeek-V3’s FP8 checkpoint: extend Chapter 39’s load_quantized to the registry, building FP8BlockLinear modules from weight and weight_scale_inv. Where: load_quantized / build in engine/formats/hf_quant.py, with registry loading in engine/models.py.
  5. ★★★ Write a Triton MLA decode kernel: one program per (request, head group), with the 576-wide rows read once for 16 or more heads. Compare with the generic kernel run with one KV head. Where: new engine/kernels/triton_mla.py, selected by MLAAttention.paged_forward in engine/models.py.

Check your understanding

  1. Why must DecoderAttention leave q_norm undefined instead of setting it to an identity module?
  2. Which RoPE pairs does Llama 3’s scaling leave alone, and why can it?
  3. Why does “dynamic NTK” scaling clash with prefix caching?
  4. How does MLA cache 576 numbers per token while giving 128 heads distinct keys and values?
  5. Why can’t the rotary part of MLA’s key be absorbed like the rest?
  6. When is the decompressed form of MLA cheaper than the absorbed form?
  7. What does group-limited routing bound, and why does that matter under expert parallelism?

Going deeper

  • Chen et al., Extending Context Window of Large Language Models via Positional Interpolation (2023); Peng et al., YaRN: Efficient Context Window Extension of Large Language Models (ICLR 2024); Meta’s Llama 3 report, The Llama 3 Herd of Models (2024), for its RoPE scaling.
  • Jiang et al., Mistral 7B (2023); Beltagy et al., Longformer (2020), for sliding-window attention.
  • DeepSeek-AI, DeepSeek-V2 (2024) for MLA and DeepSeek-V3 Technical Report (2024) for auxiliary-loss-free balancing and group-limited routing; the FlashMLA repository.
  • vLLM’s vllm/model_executor/models/ (one file per family, and registry.py) and vllm/v1/attention/backends/mla/; llama.cpp’s src/llama-model.cpp and convert_hf_to_gguf.py.