primer.ml.efficient_architectures

Long context and efficient architectures: reading a million tokens without a million-squared bill

Run: python -m primer.ml.efficient_architectures

New to the notation? primer.notation explains every symbol used here from zero. This lesson builds on attention from primer.ml.attention and the KV cache from primer.ml.inference.

Level 1: The practitioner's guide

In one sentence. Long context is paid for twice, in compute that grows with the square of the length and in cache memory that grows with the length, and every "efficient" architecture is a different bargain about which tokens a model keeps exactly, which it summarizes, and how compactly it stores them.

When you need it. You need this when a model's context window is part of your design: a whole codebase in the prompt, a day-long conversation, a document set that will not fit in retrieval, or an agent that runs for hundreds of steps. You also need it when choosing between models whose cards list different attention designs (grouped-query, sliding-window, hybrid Mamba, latent attention), because those words decide what the model will cost you per conversation and what it will forget. The tell is a context length that sounds like a solved problem. On this lesson's Llama-3-8B-shaped model, a 128,000-token conversation holds 17.2 GB of cache and a million-token one 137 GB, more than an 80 GB GPU before any weights are loaded, and the attention work grows 16,384-fold across that same range (this lesson's context-cost table). You do not need this lesson for prompts of a few thousand tokens: at that length full attention is cheap and exact, and nothing here beats it.

Your options. Each is a design you choose by choosing a model, or a setting you apply to one, from the most exact to the most compact:

Option What it does What it keeps What it costs Where it lives
Full attention with grouped-query heads Every token reads every earlier token; several query heads share each key/value head Exact recall of everything in context 128 KB per token with 8 KV heads on the 8B example, four times less than 32 heads; still 17 GB at 128k The model's architecture
Sliding-window and interleaved layers Each token reads the last w tokens; depth relays information further; some models keep a few full-attention layers Exact recall inside the window, relayed and weaker beyond it A cache capped at w tokens: 0.54 GB for a 4,096 window on the 8B example, whatever the length The model's architecture
Attention sinks with a window Keeps the first few tokens plus a recent window at serving time Fluent streaming without limit; no recall of the dropped middle Nearly nothing; StreamingLLM reports streaming to 4 million tokens and up to 22.2× speedup over recomputation The inference server
Hybrid state-space and attention Most layers carry a fixed-size state; one in several is attention Exact lookup through the attention layers, cheap everything else On the 8B shape, 4 attention layers in 32 cut the 128k cache from 17.2 GB to 2.15 GB (this lesson's hybrid formula) The model's architecture
Pure state-space or linear attention Every layer summarizes the past into a fixed state Fluent long text; blurrier exact recall A cache that never grows (4 MB for the whole 8B-shaped stack in this lesson's demo); weaker copying of exact tokens from far back The model's architecture
Latent KV cache Caches a short latent per token and rebuilds keys and values from it Exact outputs, by construction DeepSeek-V2 caches 576 numbers per token per layer instead of 32,768 (about 57× less) at the price of extra matrix work The model's architecture
A quantized KV cache Stores cached keys and values at 8 or 4 bits with one scale per vector Slightly rounded attention 3.9× smaller at 4 bits; under 1% output error at 8 bits and about 12% at 4 bits on random data (this lesson's measurement) The inference server

How to choose. Start from what the task needs to find, not from the advertised window.

  • Exact lookups across a long input (a function name in a repository, a clause in a contract): keep full attention in enough layers. A grouped-query model, or a hybrid with attention layers, and a budget for the cache.
  • Long, fluent, forward-moving text (a running transcript, a stream) with no need to quote the distant past: a sliding-window or state-space model is far cheaper, and a sink-plus-window server setting keeps even a full-attention model streaming.
  • Many concurrent long conversations on fixed hardware: the cache per token is your capacity. Prefer fewer KV heads or a latent cache, then quantize the cache, then cap the context you allow.
  • A context window claimed at 128k or more: test the model at the length you will use. RULER (Hsieh et al., 2024) found that of 17 models claiming 32k tokens or more, only half held up at 32k, and Lost in the Middle (Liu et al., 2023) found accuracy highest when the relevant passage sits at the start or the end of the input and worst in the middle.
  • Whatever you pick, put the facts the model must use where the design keeps them exactly: inside the window, near the ends, or in the prompt of a retrieval step (primer.agents.rag) instead of a million-token dump.

What it costs. Two meters run at once. Compute is the square of the length and is paid at prefill: on 4,096 tokens, full attention scores 32 times as many pairs as a 64-token window (this lesson's demo). Memory is linear in the length and is paid for as long as a conversation stays open; it decides how many users a GPU serves. The compressions differ in what they charge: grouped-query heads cost a little modelling capacity; the window costs direct access beyond it; a fixed state costs sharp recall (in this lesson's comparison, softmax attention puts up to 0.46 of a row on its favourite key where linear attention manages 0.21); the latent cache costs extra matrix work and care with positions; quantization costs precision. The gains are large: Mamba reports 5× the inference throughput of a transformer with linear scaling in length, DeepSeek-V2 reports a KV cache 93.3% smaller and 5.76× the generation throughput of its predecessor, and Jamba fits a 256k-context model on one 80 GB GPU (each paper's abstract).

What breaks.

  • The fact was outside the window. A sliding-window model can be influenced by a token 131,040 positions back (Mistral 7B's 32 layers of 4,096), but only by relay through about 25 hops, so a distant fact arrives weakened. Do not expect exact quotes from beyond the window.
  • A plain window drops the first tokens and the model falls apart. Trained models park attention on the first few tokens (attention sinks); keep them when you truncate.
  • A fixed-size state cannot copy. Pure state-space and linear-attention models do worse at reproducing a specific token from far back. If the task is retrieval or copying, keep attention layers.
  • The advertised context is not the effective context. Test at your length with your task, in the middle of the input, not only with a needle at the end.
  • A sparse pattern that saves no time. Skipped pairs only save work when the kernel skips whole blocks of the score matrix; a pattern drawn token by token costs as much as full attention.
  • A quantized cache that changes answers. Keys carry a few channels with large values; quantizing per token can lose them. Check outputs on your own data, and use a scheme like KIVI's (keys per channel, values per token) at low bit widths.

In the wild. Mistral 7B (Jiang et al.) shipped sliding-window attention with a rolling-buffer cache; Llama 3 uses grouped-query attention, the design this lesson's 8 KV heads mirror (Grattafiori et al., in primer.ml.inference). Jamba (Lieber et al.) interleaves one attention layer among Mamba layers; Mamba itself (Gu and Dao) and Mamba-2 (Dao and Gu) are the state-space models; the Sparse Transformer (Child et al.), Longformer (Beltagy et al.) and BigBird (Zaheer et al.) are the sparse patterns; StreamingLLM (Xiao et al.) is the sink-plus-window trick; DeepSeek-V2 introduced the latent cache; KIVI (Liu et al.) reaches a 2-bit cache. RULER and Lost in the Middle are the two tests to run before trusting a context length. Inference servers expose the serving-side options: vLLM and SGLang list quantization and prefix caching among their features, and the paged KV memory both use is what makes a growing cache manageable at all (primer.ml.inference).

Go deeper. Level 2 builds every row of the table from nothing: the two bills counted for real context lengths, the window mask and the reach it gives a stack of layers, sparse patterns with their pair counts, linear attention as a running sum, a state-space model that is both a loop and a convolution, Mamba's per-token step size on a recall task, the parallel scan that trains it, the latent cache with its absorbed query, and a quantized cache with its error measured. If you only needed to pick a model and size its cache, you are done.

Level 2: How it works, from scratch.

Level 2: How it works, from scratch

Attention (see primer.ml.attention) lets every token look at every earlier token. That is where its power comes from, and it is also where its bill comes from. The bill has two lines:

  • Compute. Every token is scored against every other token, so the work grows with the square of the context length n.
  • Memory. During generation the model keeps every earlier token's keys and values (the KV cache, see primer.ml.inference), so memory grows with n, per layer, per conversation.

Everything in this lesson attacks one of those two lines, in one of three ways:

  1. Look at fewer tokens: sliding-window and sparse attention.
  2. Summarize the past into a fixed-size state: linear attention and state-space models (Mamba), which run like a recurrent network.
  3. Store the past more compactly: fewer key/value heads, a small latent vector per token, fewer bits per number.
flowchart TB P["Long context is expensive<br/>compute grows with n², memory with n"] --> F["Look at fewer tokens"] P --> S["Keep a fixed-size summary"] P --> C["Store the past compactly"] F --> F1["sliding window"] & F2["sparse: local + global, strided"] S --> S1["linear attention"] & S2["state-space models, Mamba"] C --> C1["GQA / MQA"] & C2["latent KV"] & C3["quantized KV"] S2 --> H["hybrids: a few attention layers<br/>among many SSM layers"] F1 --> H

Reading it: start at the top box, the problem. The three arrows below it are the three strategies, and the leaves are the techniques this lesson builds from scratch, section by section. The bottom box is where real models often land: a mix, keeping a little full attention for exact lookups and using cheap layers for everything else.

1. Why long context is expensive

Everyday picture. A dinner party where every new guest must shake hands with every guest already there. Ten guests make 45 handshakes; a thousand guests make half a million. On top of that, the coat check keeps one coat per guest on every floor of the building, so the coat racks grow with every arrival. Attention pays both bills: the handshakes are the query-key scores, and the coats are the KV cache.

Tiny worked example. Take a Llama-3-8B-shaped model: 32 layers, 8 key/value heads, 128 numbers per head, 16-bit numbers (2 bytes). One token costs 2 × 32 × 8 × 128 × 2 = 131,072 bytes of cache, about 128 KB.

Context n Pairs scored per head per layer (n²) KV cache for one conversation
8,192 (8k) 67 million 1.1 GB
131,072 (128k) 17 billion 17.2 GB
1,048,576 (1M) 1.1 trillion 137.4 GB

Going from 8k to 1M is 128 times more tokens, 16,384 times more pairs, and a cache that no longer fits on one 80 GB GPU.

flowchart LR T["new token n"] --> Q["its query"] Q -->|"scored against"| KC[("KV cache<br/>n − 1 keys and values<br/>per layer")] KC --> O["output for token n"] T -->|"its key and value<br/>are appended"| KC

Reading it: follow one new token. Its query is scored against every key already in the cache, so the work for this one token grows with n, and the work for all n tokens grows with n². Then its own key and value are appended, so the cache grows by one entry, in every layer, for every token. The two arrows into the cache are the two bills.

Level 3: the formula and its symbols

$$ \text{pairs} = n^2 \qquad\qquad \text{KV bytes} = 2 \, L \, H_{kv} \, d_h \, b \, n $$

Symbols

Symbol Meaning here In the example
$n$ context length: how many tokens are in play 131,072
$n^2$ $n$ times $n$: every query paired with every key (causal masking halves it, which does not change how it grows) 17,179,869,184
$2$ one key and one value per token
$L$ layers, each with its own cache 32
$H_{kv}$ key/value heads per layer 8
$d_h$ numbers per head 128
$b$ bytes per stored number 2 (16-bit)

In words: "the score work is the context length squared, and the cache holds a key and a value for every head, in every layer, for every token."

With the numbers: at n = 131,072 the pairs are 131,072² ≈ 1.7 × 10¹⁰, and the cache is 2 × 32 × 8 × 128 × 2 × 131,072 = 17,179,869,184 bytes ≈ 17.2 GB.

Level 3: in Python

In Python:

L, H_kv, d_h, b = 32, 8, 128, 2
# bytes per token: a key and a value, per layer, per KV head
per_token = 2 * L * H_kv * d_h * b
per_token  # → 131072
contexts = [8_192, 131_072, 1_048_576]
# pairs per head per layer
[f"{n * n:,}" for n in contexts]  # → ['67,108,864', '17,179,869,184', '1,099,511,627,776']
# GB of cache for one conversation
[round(n * per_token / 1e9, 1) for n in contexts]  # → [1.1, 17.2, 137.4]

Pairs scored grow from 67 million at 8k tokens to 1.1 trillion at 1M, and the KV cache from 1.1 GB to 137 GB, past an 80 GB GPU

Reading it: the left panel counts query-key pairs per head per layer, on a logarithmic axis where each gridline is ten times the last: every step to the right multiplies the bar by 256 and then by 64. The right panel is the KV cache for one conversation on an ordinary axis, with the dashed line at 80 GB, the memory of one large GPU. The 1M bar crosses it before any weights are loaded. The left panel is the compute bill; the right panel is the memory bill.

Why it matters in practice. The compute bill is paid once per prompt (prefill), and exact tricks such as FlashAttention make it faster without changing it. The memory bill is paid for as long as a conversation is open, and it decides how many users one GPU can serve. That is why most of the techniques below target the cache.

In code: context_cost returns both lines of the bill for any context length, using primer.ml.inference.kv_cache_bytes for the cache.

2. Sliding-window attention: reading through a letterbox

Everyday picture. You read a long scroll through a letterbox slot that shows only the last w words. On your own you would lose the beginning. But suppose a row of readers sits one above the other, each reading the notes of the reader below through a slot of the same width. Each reader passes things along a little further, like a bucket brigade, so a stack of readers can carry a fact much further back than any one slot shows.

Tiny worked example. Eight tokens, window w = 3: each token reads itself and the two tokens before it.

            key:  0 1 2 3 4 5 6 7
query 0           ✓ · · · · · · ·
query 1           ✓ ✓ · · · · · ·
query 2           ✓ ✓ ✓ · · · · ·
query 3           · ✓ ✓ ✓ · · · ·
query 4           · · ✓ ✓ ✓ · · ·
query 5           · · · ✓ ✓ ✓ · ·
query 6           · · · · ✓ ✓ ✓ ·
query 7           · · · · · ✓ ✓ ✓

That is 1 + 2 + 3 × 6 = 21 scores instead of the 36 a full causal mask needs. Now stack three such layers. Token 7 reads token 5 in the top layer; token 5 had read token 3 in the layer below; token 3 had read token 1 in the layer below that. So after three layers, token 1 can influence token 7, six positions back, but token 0 cannot.

flowchart RL subgraph L3["layer 3"] a7["token 7"] end subgraph L2["layer 2"] b5["token 5"] end subgraph L1["layer 1"] c3["token 3"] end subgraph IN["input"] d1["token 1"] d0["token 0: out of reach"] end a7 -->|"reads 2 back"| b5 -->|"reads 2 back"| c3 -->|"reads 2 back"| d1

Reading it: read right to left, from token 7 at the top layer down to the input. Each arrow is one layer's window: it can reach at most w − 1 = 2 positions back. Three layers chain three arrows, so the furthest input token 7 can hear from is 3 × 2 = 6 positions back, token 1. Token 0 would need a fourth hop. Information still travels far, but only in stages.

Level 3: the formula and its symbols

$$ M_{ij} = \begin{cases} 1 & \text{if } 0 \le i - j < w \ 0 & \text{otherwise} \end{cases} \qquad\qquad \text{reach} = L\,(w - 1) $$

Symbols

Symbol Meaning here In the example
$M_{ij}$ the mask: 1 if query $i$ may read key $j$, 0 if not row 5: keys 3, 4, 5
$i$, $j$ the query's position and the key's position $i = 5$
$i - j$ how far back the key is; negative means the future 0, 1, 2 allowed
$w$ the window: how many tokens each query reads, itself included 3
${$ "cases": use the line whose condition holds
$L$ how many windowed layers are stacked 3
reach the furthest back an input can influence an output 6

In words: "a query may read a key if the key is not in the future and is fewer than w positions back; stacking L layers lets information travel L times w − 1 positions."

With the numbers: row 5 allows j = 3, 4, 5 (5 − 3 = 2 < 3). Three layers of w = 3 reach 3 × 2 = 6. Mistral 7B uses w = 4,096 over 32 layers: 32 × 4,095 = 131,040 tokens, about 131k, from a window of 4k.

Level 3: in Python

In Python:

n, w = 8, 3
# M_ij for the row of token 5
[int(0 <= 5 - j < w) for j in range(n)]  # → [0, 0, 0, 1, 1, 1, 0, 0]
# pairs scored: row i keeps min(i + 1, w) keys
sum(min(i + 1, w) for i in range(n))  # → 21
# full causal attention, for comparison
n * (n + 1) // 2  # → 36
# reach after L layers
L = 3
L * (w - 1)  # → 6
# Mistral 7B: 32 layers, a 4,096-token window
32 * (4096 - 1)  # → 131040

Left, one layer with a window of 4 is a thin diagonal band; right, after three layers each token can hear the last 10 positions, a band three times as wide

Reading it: both panels are 16 × 16 grids, with rows for the token doing the looking and columns for the token being looked at, like the causal heatmap in primer.ml.attention. Dark means "can reach". On the left, one layer with w = 4 is a thin band along the diagonal: 4 cells per row at most. On the right, the same mask applied three times: the band has widened to 3 × 3 + 1 = 10 cells, because each layer extends the reach by w − 1 = 3. Neither panel ever touches the upper-right triangle: the window is still causal.

In code: sliding_window_mask builds the band, pairs_computed counts its cells, receptive_field applies the mask layer after layer to find who can reach whom, and reach is the L(w − 1) formula.

The rolling buffer: a cache that stops growing

Everyday picture. A whiteboard with room for exactly w notes. When it is full, the next note goes over the oldest one. You never need a bigger board, however long the meeting runs.

Tiny worked example. With w = 4, after token 9 the buffer holds the keys and values of tokens 6, 7, 8 and 9. Token 10 overwrites token 6's slot. On the running 8B example at 128k tokens, a 4,096-token window holds 4,096 × 128 KB = 0.54 GB instead of 17.2 GB: 32 times less, and the same at 1M tokens.

flowchart LR subgraph B["buffer of w = 4 slots, after token 9"] s0["slot 0: token 8"] s1["slot 1: token 9"] s2["slot 2: token 6 (oldest)"] s3["slot 3: token 7"] end N["token 10 arrives"] -->|"overwrites the oldest"| s2 B --> A["token 10 attends to<br/>tokens 7, 8, 9, 10"]

Reading it: the four slots are all the memory there is. Slot number is token number modulo 4 (the remainder after dividing by 4), so token 10 lands in slot 2, where token 6 sat. Token 6 was about to leave the window anyway, so nothing the window needs is lost.

Why it matters in practice. Mistral 7B combined the window with this rolling buffer. Several later model families interleave sliding-window layers with a few full-attention layers: the local layers are cheap, and the full layers keep long-range lookups exact. The weakness is plain from the diagram: a fact beyond the reach is invisible, and a fact inside the reach has to survive several hops.

In code: sliding_window_decode generates token by token with a buffer that never holds more than w keys, and returns the same outputs as masked attention with sliding_window_mask.

3. Sparse attention: a few long-distance lines

Everyday picture. An open-plan office. You mostly talk to the people at the desks next to yours (local). Anyone can phone the front desk, and the front desk hears from everyone (a global token): any two people are at most two calls apart. Another design is the express train: you talk to your neighbours and also to every fourth desk down the row (strided), so a message can travel far in a few big jumps.

Tiny worked example. Sixteen tokens, window 4.

  • The window alone: 1 + 2 + 3 + 4 × 13 = 58 pairs.
  • Make token 0 global: the 12 rows past the window (tokens 4 to 15) add a pair each: 70 pairs, against 136 for full causal attention.
  • Strided with stride 4: the last 4 tokens, plus every 4th token before them. Token 13 reads 10, 11, 12, 13 and then 9, 5, 1. Over all rows that is 58 + 24 = 82 pairs.
flowchart LR G(("token 0<br/>global")) t3["token 3"] --- G t7["token 7"] --- G t11["token 11"] --- G t15["token 15"] --- G t14["token 14"] --- t15 t13["token 13"] --- t14

Reading it: the circle is the global token. Every token has a line to it, so token 15 can reach token 3 in two hops, through token 0, however long the sequence. The short lines at the bottom are the ordinary local window. The picture is sparse (few lines), yet no two tokens are far apart.

Level 3: the formula and its symbols

$$ M_{ij} = 1 \quad\text{when}\quad i - j \ge 0 \;\text{ and }\; \big(\, i - j < s \;\text{ or }\; (i - j) \bmod s = 0 \,\big) $$

Symbols

Symbol Meaning here In the example
$M_{ij}$ 1 if query $i$ may read key $j$
$i - j$ how far back key $j$ is for $i = 13$: 0 to 13
$s$ the stride: the local width, and the jump between long-range keys 4
$\bmod$ "modulo": the remainder after dividing; $(i - j) \bmod s = 0$ means "a whole number of strides back" 8 mod 4 = 0
and, or both conditions must hold; at least one must hold

In words: "never read the future; read the last s tokens, and beyond them every s-th token."

With the numbers: for i = 13, s = 4: distances 0 to 3 give keys 13, 12, 11, 10; distances 4, 8 and 12 give keys 9, 5, 1.

Level 3: in Python

In Python:

s = 4
# row 13: which keys does token 13 read?
[j for j in range(16) if 13 - j >= 0 and (13 - j < s or (13 - j) % s == 0)]  # → [1, 5, 9, 10, 11, 12, 13]
# strided pairs over all 16 rows
sum(1 for i in range(16) for j in range(i + 1) if i - j < s or (i - j) % s == 0)  # → 82
# a window of 4 plus global token 0
sum(1 for i in range(16) for j in range(i + 1) if i - j < 4 or j == 0)  # → 70

Four 16 by 16 masks: full causal attention scores 136 pairs, the window 58, window plus a global first token 70, strided 82

Reading it: each panel is a mask, rows for queries and columns for keys, dark where a score is computed; the title gives the count. The first panel is full causal attention, a solid triangle. The window is a diagonal band. Adding a global token fills in the first column: everyone reads token 0. The strided pattern adds dotted diagonals every 4 columns: the express stops. At 16 tokens the savings look modest; with a fixed window they grow with n, because the band stays w wide while the triangle keeps growing.

Why it matters in practice. The Sparse Transformer introduced strided patterns; Longformer and BigBird combined windows with global tokens (BigBird added a few random links too) to read documents of thousands of tokens. StreamingLLM found that trained models park a lot of attention on the very first tokens (attention sinks); a plain window drops them and falls apart, while keeping the first few tokens plus a window lets a model stream indefinitely. One caution: a sparse pattern only saves time if the GPU kernel skips whole blocks of the score matrix, which is why real patterns are built from blocks.

In code: global_local_mask adds global rows and columns to a window, strided_mask builds the express-stop pattern, and pairs_computed counts what each one scores. Any of them can be passed as the mask to primer.ml.attention.scaled_dot_product_attention, and every skipped pair gets a weight of exactly 0.

4. Linear attention: a pot instead of a guest list

Everyday picture. A potluck soup. In softmax attention every new guest tastes every dish on the table, one by one, and then mixes a bowl: the more dishes, the longer it takes. In linear attention every guest pours their dish into one shared pot as they arrive, and a new guest takes a single ladle, seasoned to their own taste. The pot never grows, and the ladle costs the same for guest 3 as for guest 3 million.

The trick that makes the pot possible: softmax's score $e^{q \cdot k}$ cannot be split into "a part that depends on q" times "a part that depends on k". Replace it with $\phi(q) \cdot \phi(k)$, a dot product of transformed vectors, and it can. Then the order of the matrix multiplies can be swapped: $(\phi(Q)\phi(K)^\top)\,V = \phi(Q)\,(\phi(K)^\top V)$. The left side builds an n × n matrix; the right side builds a small d × d one (the pot) and never builds the big one. Matrix multiplication allows regrouping like this (associativity), which primer.notation covers under matrix multiply.

Tiny worked example. Three tokens with 2-number keys (1, 0), (0, 1), (1, 1) and one-number values 2, 4, 6. The third token's query is (1, 0). Use the feature map φ(x) = elu(x) + 1, which for a positive number is simply x + 1 and for zero or a negative number is $e^x$ (always above zero).

  1. φ of the keys: (2, 1), (1, 2), (2, 2).
  2. Pour into the pot: S = (2, 1)·2 + (1, 2)·4 + (2, 2)·6 = (20, 22), and the running total of keys z = (5, 5).
  3. Ladle with φ(q) = (2, 1): (2·20 + 1·22) / (2·5 + 1·5) = 62 / 15 = 4.13.

The slow way agrees: the weights φ(q)·φ(k) are 5, 4 and 6, so the output is (5·2 + 4·6 + 6·6) / 15 = 62/15. Softmax attention on the same numbers gives 4.00: linear attention is a different attention, not a faster copy of the same one.

flowchart LR subgraph T["each new token i"] K["φ(k_i)"] V["v_i"] Qi["φ(q_i)"] end K & V -->|"add φ(k_i) v_iᵀ"| S[("pot S<br/>d_k × d_v numbers")] K -->|"add φ(k_i)"| Z[("total z<br/>d_k numbers")] Qi --> R["output = φ(q_i)ᵀ S / φ(q_i)ᵀ z"] S --> R Z --> R

Reading it: each token does two things. Its key and value go into the pot S (and its key into the running total z, used to normalize). Its query reads the pot once. S and z have a fixed size set by the head width, not by the number of tokens, so this is a recurrent network: a state updated once per token, with no cache that grows.

Level 3: the formula and its symbols

$$ o_i = \frac{\sum_{j \le i} \big(\phi(q_i) \cdot \phi(k_j)\big)\, v_j}{\sum_{j \le i} \phi(q_i) \cdot \phi(k_j)} = \frac{\phi(q_i)^\top S_i}{\phi(q_i)^\top z_i}, \qquad S_i = S_{i-1} + \phi(k_i)\, v_i^\top, \qquad z_i = z_{i-1} + \phi(k_i) $$

Symbols

Symbol Meaning here In the example
$o_i$ the output for token $i$ $o_3 = 4.13$
$q_i$, $k_j$, $v_j$ query of token $i$; key and value of token $j$ $q_3 = (1, 0)$
$\phi$ the feature map, applied to each number: $x + 1$ if $x > 0$, else $e^x$; always positive so no weight is negative φ(1, 0) = (2, 1)
$\sum_{j \le i}$ add up over every token $j$ up to and including $i$ (causal) $j$ = 1, 2, 3
$\phi(q_i) \cdot \phi(k_j)$ the unnormalized weight, a dot product in place of $e^{q \cdot k}$ 5, 4, 6
$v_i^\top$ the value laid on its side as a row
$\phi(k_i)\, v_i^\top$ an outer product: a column times a row, giving a small table whose entry (m, c) is $\phi(k_i)_m \, v_{i,c}$ (2, 1) × 2 = (4, 2)
$S_i$ the pot after token $i$: the sum of those tables (20, 22)
$z_i$ the running sum of $\phi(k_j)$, for the denominator (5, 5)
$\phi(q_i)^\top S_i$ the query's ladle: its dot product with each column of $S_i$ 62

In words: "each output is a weighted average of the values so far, with weights φ(q)·φ(k); because those weights split into a query part and a key part, the key-and-value part can be kept as a running sum, and each query reads that sum once."

With the numbers: S₃ = (20, 22), z₃ = (5, 5), φ(q₃) = (2, 1): o₃ = (40 + 22) / (10 + 5) = 62/15 = 4.133.

Level 3: in Python

In Python:

import math
def phi(v):
    # elu(x) + 1: x + 1 above zero, e^x at or below it
    return [x + 1 if x > 0 else math.exp(x) for x in v]
def dot(a, b):
    return sum(a_m * b_m for a_m, b_m in zip(a, b))
keys = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]
values = [2.0, 4.0, 6.0]
q3 = [1.0, 0.0]
[phi(k) for k in keys]  # → [[2.0, 1.0], [1.0, 2.0], [2.0, 2.0]]
# S: the running sum of φ(k_j) v_j; z: the running sum of φ(k_j)
S = [sum(phi(k)[m] * v for k, v in zip(keys, values)) for m in range(2)]
S  # → [20.0, 22.0]
z = [sum(phi(k)[m] for k in keys) for m in range(2)]
z  # → [5.0, 5.0]
# o_3 = φ(q)ᵀS / φ(q)ᵀz
round(dot(phi(q3), S) / dot(phi(q3), z), 3)  # → 4.133
# the quadratic form: the weights φ(q)·φ(k_j), then their weighted average
weights = [dot(phi(q3), phi(k)) for k in keys]
weights  # → [5.0, 4.0, 6.0]
round(dot(weights, values) / sum(weights), 3)  # → 4.133
# softmax attention on the same numbers (unscaled), for contrast
e = [math.exp(dot(q3, k)) for k in keys]
round(dot(e, values) / sum(e), 3)  # → 4.0

Side by side on the same queries and keys: softmax attention puts almost half of each late row on one key, while linear attention spreads each row thinly, its largest weight about 0.2

Reading it: the two heatmaps use the same ten queries and keys; rows are queries, columns are keys, and darker means more weight. Both are causal triangles and every row sums to 1. On the left, softmax picks favourites: the exponential stretches gaps between scores (see primer.ml.attention), so in the last five rows the largest weight averages 0.46. On the right, linear attention spreads the same rows thinly: its largest weight averages 0.21. The queries here are twice the usual size, and that is the telling part: at the usual size the two numbers are 0.28 and 0.20, so making a query more decisive sharpens softmax a lot and linear attention hardly at all, because φ(q)·φ(k) never stretches a gap the way $e^x$ does. That bluntness is the price of the pot: a fixed-size summary cannot pick out one exact token as sharply.

Why it matters in practice. The work drops from n² × d to n × d², and generation needs only the pot, not a cache. The catch is recall: asked to copy back a specific token from far away, a fixed-size pot does worse than a full cache. Later work added forgetting (decay) to the pot, which turns out to be exactly the state-space models of the next section; the Mamba-2 paper shows the two views describe the same computation.

In code: feature_map is φ, linear_attention_quadratic builds all n × n weights (from linear_attention_weights) and blends the values, and linear_attention_recurrent gets the same outputs from the two running sums, returning the fixed-size S and z.

5. State-space models: a summary with a fixed update rule

A recurrence that is also a convolution

Everyday picture. A cup of tea cooling on a desk. Every minute it keeps some fraction of its heat and gains whatever hot water you pour in. Its temperature right now is a summary of everything you ever poured, with older pours counting less. That is a recurrent network's one-page summary (see primer.ml.cnn_rnn), with one difference: the rewrite rule is linear (multiply and add, nothing more), and linearity buys a second way to compute the same thing.

Tiny worked example. The state keeps half of itself each step: A = 0.5, B = 1, C = 1. Inputs x = (1, 0, 0, 2).

step t x_t h_t = 0.5 · h_(t−1) + x_t y_t
1 1 0.5 · 0 + 1 = 1 1
2 0 0.5 · 1 + 0 = 0.5 0.5
3 0 0.5 · 0.5 + 0 = 0.25 0.25
4 2 0.5 · 0.25 + 2 = 2.125 2.125

Unroll the loop and each output is a weighted sum of all past inputs with weights 1, 0.5, 0.25, 0.125 for "now, 1 step ago, 2 steps ago, 3 steps ago": y₄ = 1·2 + 0.5·0 + 0.25·0 + 0.125·1 = 2.125. The same answer, with no loop.

flowchart LR X["inputs x_1 ... x_T"] --> R["recurrence<br/>h_t = A h_(t−1) + B x_t<br/>one step at a time"] X --> K["convolution<br/>slide the kernel (CB, CAB, CA²B, ...)<br/>all positions at once"] R --> Y["the same outputs y_1 ... y_T"] K --> Y

Reading it: two roads from the same inputs to the same outputs. The top road is how the model generates: one token at a time, carrying a small state h, with constant memory. The bottom road is how it trains: the kernel is computed once, and every output is a weighted sum that can be done in parallel, like a convolution in primer.ml.cnn_rnn. Classic RNNs only had the top road, which is why they trained slowly.

Level 3: the formula and its symbols

$$ h_t = A\,h_{t-1} + B\,x_t, \qquad y_t = C\,h_t \qquad\Longleftrightarrow\qquad y_t = \sum_{k=0}^{t-1} \big(C A^{k} B\big)\, x_{t-k} $$

Symbols

Symbol Meaning here In the example
$x_t$ the input at step $t$ (one channel) (1, 0, 0, 2)
$h_t$ the state after step $t$: $N$ numbers (here $N$ = 1); $h_0 = 0$ 1, 0.5, 0.25, 2.125
$A$ $N \times N$ matrix: how the old state carries over (its size sets how fast things fade) 0.5
$B$ how the input is written into the state 1
$C$ how the state is read out 1
$y_t$ the output at step $t$
$\Longleftrightarrow$ "the same thing, written another way"
$A^k$ $A$ multiplied by itself $k$ times; $A^0$ is "do nothing" 0.5³ = 0.125
$C A^k B$ the kernel: how much an input $k$ steps ago still counts now 1, 0.5, 0.25, 0.125
$\sum_{k=0}^{t-1}$ add over every look-back distance $k$ from 0 to $t - 1$

In words: "the state keeps a fraction of itself and adds the new input, and the output reads the state; equivalently, the output is every past input weighted by how much it has faded since."

With the numbers: y₄ = C A⁰ B x₄ + C A¹ B x₃ + C A² B x₂ + C A³ B x₁ = 1·2 + 0.5·0 + 0.25·0 + 0.125·1 = 2.125.

Level 3: in Python

In Python:

A, B, C = 0.5, 1.0, 1.0
x = [1.0, 0.0, 0.0, 2.0]
h, y = 0.0, []
for x_t in x:
    # keep half of the old state, add the new input
    h = A * h + B * x_t
    y.append(C * h)
y  # → [1.0, 0.5, 0.25, 2.125]
# the kernel C A^k B, for k = 0, 1, 2, 3
kernel = [C * A ** k * B for k in range(4)]
kernel  # → [1.0, 0.5, 0.25, 0.125]
# y_4 as a weighted sum of all four inputs, newest first
sum(kernel[k] * x[3 - k] for k in range(4))  # → 2.125

Three kernels: with A = 0.5 an input is forgotten within about 5 steps, with 0.9 within about 40, with 0.99 it still counts about a fifth after 150 steps

Reading it: the x-axis is how many steps ago an input arrived; the y-axis is how much it still counts (the kernel C A^k B). Each curve is one value of A. With A = 0.5 the curve collapses almost at once: short memory. With A = 0.99 it is still well above zero after 150 steps: long memory. A real SSM layer has thousands of channels, each with its own A, so it remembers at many time scales at once. S4 was the first to make this work on very long sequences, by choosing and computing these kernels carefully.

In code: ssm_recurrent runs the loop, ssm_kernel builds C A^k B, and ssm_convolution gets the same outputs from the kernel with a single NumPy convolution.

Selective: letting each token decide what to keep (Mamba)

Everyday picture. A note-taker with a dial. For filler words ("um", "so", "anyway") they barely touch their notes. For a name or a number they wipe the relevant line and write the new fact. A fixed SSM uses the same dial setting for every word, so it must either write everything (and forget quickly) or write little (and never take in the important word properly). Mamba reads the dial setting off each token itself: that is what selective means.

The dial is a step size Δ. The parameters of the state update are derived from it at every step (this is called discretization: turning a continuous rate of change into one step's keep and write amounts):

  • keep factor $\bar{A} = e^{\Delta a}$, with $a$ negative, so a big Δ makes it nearly 0 (forget) and a tiny Δ nearly 1 (keep);
  • write factor $\bar{B}$, which goes the opposite way.

Tiny worked example. One channel, a = −1, b = 1. A token that carries a marker gets Δ ≈ 5; an ordinary token gets Δ ≈ 0.0067.

token Δ keep Ā = e^(−Δ) write B̄ = 1 − e^(−Δ) effect
marked 5.007 0.0067 0.9933 replace the state with this token
ordinary 0.0067 0.9933 0.0067 leave the state almost untouched

The recall task: a sequence of small noise values with one marked 7 in third place, then nine more noise values. The selective SSM writes the 7 almost fully and then keeps it: at the end its state is 6.55. With one fixed Δ for every token, the best any setting manages is 0.26: a Δ big enough to write the 7 also lets the nine later tokens overwrite it.

flowchart LR U["token u_t"] --> D["Δ_t = softplus(w · u_t + β)<br/>how much this token matters"] D --> AB["Ā_t = e^(Δ_t a): keep<br/>B̄_t: write"] U --> XV["x_t: what to write"] AB --> H["h_t = Ā_t h_(t−1) + B̄_t x_t"] XV --> H HP["h_(t−1)"] --> H H --> Y["y_t = C h_t"]

Reading it: the new part is the top path. The token itself feeds a small function that produces its step size Δ, which sets how much of the old state to keep and how much of this token to write. In Mamba the matrices B and C are also computed from the token, by the same kind of path. Everything below is the recurrence from before, except that Ā and B̄ now change at every step.

Level 3: the formula and its symbols

$$ \Delta_t = \operatorname{softplus}(w\,u_t + \beta), \qquad \bar{A}_t = e^{\Delta_t a}, \qquad \bar{B}_t = \frac{e^{\Delta_t a} - 1}{a}\, b, \qquad h_t = \bar{A}_t\, h_{t-1} + \bar{B}_t\, x_t $$

Symbols

Symbol Meaning here In the example
$u_t$ what the model can see about token $t$ (here, just its marker, 0 or 1) 1 for the 7
$w$, $\beta$ learned weight and offset that turn $u_t$ into a step size 10, −5
softplus $\log(1 + e^x)$: a smooth ramp that is always positive, about $x$ for large $x$ and about 0 for very negative $x$ softplus(5) = 5.007
$\Delta_t$ the step size for token $t$: how much this token matters 5.007 or 0.0067
$a$ a negative learned rate; more negative means faster forgetting −1
$b$ how strongly inputs are written 1
$\bar{A}_t$ this step's keep factor ("A-bar") 0.0067 or 0.9933
$\bar{B}_t$ this step's write factor ("B-bar") 0.9933 or 0.0067
$x_t$, $h_t$ the value written, and the state after step $t$ 7, then 6.95

In words: "each token computes how much it matters; that sets how much of the old state survives and how much of the token gets written; then the usual update runs with those per-token amounts."

With the numbers: the marked 7 has Δ = softplus(10·1 − 5) = 5.007, so Ā = e^(−5.007) = 0.0067 and B̄ = (0.0067 − 1)/(−1) · 1 = 0.9933: the state becomes about 0.9933 × 7 = 6.95. Each ordinary token after it keeps 0.9933 of the state, and nine of them leave about 6.95 × 0.9933⁹ = 6.55.

Level 3: in Python

In Python:

import math
def softplus(x):
    return math.log(1 + math.exp(x))
a, b = -1.0, 1.0
# Δ for a marked token, then for an ordinary one
d_mark, d_plain = softplus(10 * 1 - 5), softplus(10 * 0 - 5)
round(d_mark, 3), round(d_plain, 4)  # → (5.007, 0.0067)
# Ā_t and B̄_t for the marked token: keep almost nothing, write almost everything
keep_mark, write_mark = math.exp(d_mark * a), (math.exp(d_mark * a) - 1) / a * b
round(keep_mark, 4), round(write_mark, 4)  # → (0.0067, 0.9933)
# and for an ordinary token: keep almost everything, write almost nothing
keep_plain = math.exp(d_plain * a)
round(keep_plain, 4)  # → 0.9933
# the 7 is written once, then kept through nine ordinary tokens
round(write_mark * 7 * keep_plain ** 9, 2)  # → 6.55

The selective state jumps to about 7 at the marked token and holds near 6.5 to the end, while the best fixed step peaks near 0.6 and ends at 0.26, and a large fixed step just copies the latest noise

Reading it: the x-axis is the position in the sequence; grey bars are the input values, with the marked 7 at position 2. The blue line is the selective SSM's state: it jumps to the 7 and then barely moves, because every later token has a tiny Δ. The red line is the best fixed Δ: it can only write a small fraction of each input, so the 7 lifts it to about 0.6 at most and it ends at 0.26. The orange line is a large fixed Δ: it writes every token fully, so it tracks whatever arrived last and the 7 is gone by the next step. Only the selective model can both write the important token and ignore the rest.

In code: discretize turns Δ into Ā and B̄, softplus and step_sizes compute each token's Δ, selective_ssm runs the per-token recurrence, and selective_recall and time_invariant_recall run the recall task.

Training in parallel when the kernel keeps changing: the scan

Everyday picture. A relay race where each runner must know the total time so far. Done one after another it takes as long as the whole race. But "multiply by a, then add b" steps can be merged: two consecutive steps are themselves one step of the same shape. So pairs of runners merge their legs, then pairs of pairs, like a knockout tournament: log₂ n rounds instead of n (log₂ n is how many times n can be halved before reaching 1: 10 for 1,024).

Tiny worked example. The four steps of the tea example are (a, b) = (0.5, 1), (0.5, 0), (0.5, 0), (0.5, 2), meaning "h becomes a·h + b". Round 1: every step merges with the one before it. Round 2: every step merges with the result two places before it. After 2 rounds (log₂ 4 = 2), the b parts are 1, 0.5, 0.25, 2.125: every state from the loop.

flowchart TB s1["(0.5, 1)"] s2["(0.5, 0)"] s3["(0.5, 0)"] s4["(0.5, 2)"] s1 --> r2["(0.25, 0.5)"] s2 --> r2 s2 --> r3["(0.25, 0)"] s3 --> r3 s3 --> r4["(0.25, 2)"] s4 --> r4 s1 --> f3["(0.125, 0.25)"] r3 --> f3 r2 --> f4["(0.0625, 2.125)"] r4 --> f4

Reading it: the top row holds the four steps. Each arrow pair is one merge; every merge in a row happens at the same time. The middle row is round 1, the bottom row round 2. Read the second number in each final box (and in s1 and r2, which were already complete): 1, 0.5, 0.25, 2.125, the states from the table above. A selective SSM cannot use the convolution road, because its Ā changes every step, but it can use this one.

Level 3: the formula and its symbols

$$ (a_1, b_1) \circ (a_2, b_2) = \big(a_1 a_2,\; a_2 b_1 + b_2\big) $$

Symbols

Symbol Meaning here In the example
$(a, b)$ one step: "multiply the state by $a$, then add $b$" (0.5, 1)
$\circ$ "do the first step, then the second": merging two steps into one
$a_1 a_2$ the combined multiplier 0.25
$a_2 b_1 + b_2$ the first step's addition, shrunk by the second step, plus the second's own 0.5·1 + 0 = 0.5

In words: "doing step 1 then step 2 is the same as one step that multiplies by both and adds step 1's contribution, faded by step 2."

With the numbers: (0.25, 0.5) ∘ (0.25, 2) = (0.0625, 0.25·0.5 + 2) = (0.0625, 2.125), the last box in the diagram.

Level 3: in Python

In Python:

def merge(first, then):
    a1, b1 = first
    a2, b2 = then
    return (a1 * a2, a2 * b1 + b2)
steps = [(0.5, 1.0), (0.5, 0.0), (0.5, 0.0), (0.5, 2.0)]
# round 1: each step absorbs the one just before it
r1 = [steps[0]] + [merge(steps[t - 1], steps[t]) for t in range(1, 4)]
r1  # → [(0.5, 1.0), (0.25, 0.5), (0.25, 0.0), (0.25, 2.0)]
# round 2: each absorbs the result two places before it
r2 = r1[:2] + [merge(r1[t - 2], r1[t]) for t in range(2, 4)]
[b for a, b in r2]  # → [1.0, 0.5, 0.25, 2.125]

Why it matters in practice. This is how Mamba trains on long sequences at GPU speed despite being recurrent: a parallel scan, written so the state stays in fast on-chip memory. At generation time it switches back to the plain loop, one token at a time, with a state that never grows.

In code: parallel_scan runs the rounds for any length and returns the states and the number of rounds: 10 for 1,024 steps.

Constant memory, and hybrids

Everyday picture. An SSM travels with a backpack of fixed size; the KV cache is a suitcase that grows with every token. A hybrid model mostly uses backpacks but brings one suitcase for every few layers, so it can still look up an exact earlier token when it needs to.

Tiny worked example. One SSM layer with 4,096 channels and a 16-number state per channel holds 4,096 × 16 × 2 bytes = 128 KB, at 10 tokens or at 10 million. One attention layer of the running 8B example holds 537 MB at 128k tokens. Build 32 layers as 4 attention layers and 28 SSM layers (one in eight, the ratio Jamba uses) and the cache at 128k drops from 17.2 GB to 2.15 GB.

flowchart TB I["tokens in"] --> M1["SSM layer<br/>fixed state"] --> M2["SSM layer"] --> M3["..."] --> A1["attention layer<br/>KV cache grows"] A1 --> M4["SSM layer"] --> M5["..."] --> A2["attention layer"] --> O["next-token scores"]

Reading it: most boxes are SSM layers, each carrying a fixed-size state. Every so often an attention layer sits in the stack; only those keep a KV cache. The memory bill therefore scales with the number of attention layers, not with the total depth, while the attention layers are still there for exact copying and lookup.

Level 3: the formula and its symbols

$$ \text{bytes} = L_{\text{att}} \cdot n \cdot 2\,H_{kv}\,d_h\,b \;+\; (L - L_{\text{att}}) \cdot D \cdot N \cdot b $$

Symbols

Symbol Meaning here In the example
$L$ total layers 32
$L_{\text{att}}$ how many of them are attention layers 4
$n$ context length 131,072
$2\,H_{kv}\,d_h\,b$ one attention layer's cache per token (from section 1) 4,096 bytes
$D$ channels in an SSM layer 4,096
$N$ state numbers per channel 16
$b$ bytes per number 2

In words: "the attention layers pay per token as before; the SSM layers pay a fixed amount that does not depend on n at all."

With the numbers: 4 × 131,072 × 4,096 = 2,147,483,648 bytes, plus 28 × 4,096 × 16 × 2 = 3,670,016 bytes: 2.15 GB, about one eighth of 17.2 GB.

Level 3: in Python

In Python:

n, H_kv, d_h, b = 131_072, 8, 128, 2
D, N = 4096, 16
# one SSM layer's state, in KB, at any context length
D * N * b / 1024  # → 128.0
# one attention layer's KV cache at 128k tokens, in MB
n * 2 * H_kv * d_h * b / 1e6  # → 536.870912
# 32 layers: all attention, then 4 attention + 28 SSM, in GB
round(32 * n * 2 * H_kv * d_h * b / 1e9, 2)  # → 17.18
round((4 * n * 2 * H_kv * d_h * b + 28 * D * N * b) / 1e9, 2)  # → 2.15

Why it matters in practice. Pure SSMs are strong at language modelling but weaker than attention at copying and exact recall over long contexts, for the same reason as linear attention: a fixed-size state is a summary. Hybrids such as Jamba keep a few attention layers for those jobs and get most of the SSM's memory savings.

In code: ssm_state_bytes is the fixed backpack, and hybrid_cache_bytes adds up a mixed stack.

6. Compressing the KV cache

Fewer key/value heads: GQA and MQA, a recap

Everyday picture. Colleagues sharing one reference binder instead of each keeping a personal copy: everyone still asks their own questions, but the shelf holds fewer binders.

Tiny worked example. On the running example (32 layers, 128 numbers per head, 16-bit), the cache per token for different numbers of KV heads:

Design KV heads Cache per token
multi-head attention 32 512 KB
grouped-query attention 8 128 KB
multi-query attention 1 16 KB

The diagram of shared heads and the full memory arithmetic live in primer.ml.attention and primer.ml.inference; everything below stacks on top of whichever of these a model uses.

Latent KV: cache the ingredients, cook on demand

Everyday picture. A restaurant that stores ingredients, not finished dishes. Every head's keys and values can be cooked from a short list of ingredients (the latent vector) with a fixed recipe shared by every token. Store the ingredients, keep the recipe once, and cook when a query needs it. Better still, the recipe can be folded into the query itself, so nothing is ever cooked at all.

Tiny worked example. A token x = (1, 0, 2, 1), two heads of width 2. Ordinary attention would cache 2 heads × 2 numbers × (key and value) = 8 numbers. Instead:

  1. Squeeze: c = x W_down = (1 + 2, 0 + 1) = (3, 1). Only these 2 numbers are cached: 4 times smaller.
  2. When needed, expand: k = c W_uk = (3, 1, 3, 1), both heads' keys.
  3. Or fold the expansion into the query: for q = (1, 2, 0, 1), q · k = 6, and (q W_ukᵀ) · c = (1, 3) · (3, 1) = 6. Same score, and k was never built.

DeepSeek-V2 uses this at scale: its 128 heads of width 128 would cache 32,768 numbers per token per layer; its latent caches 512, plus 64 for a small key that carries position information: 576, about 57 times less.

flowchart LR X["token x<br/>(d_model numbers)"] -->|"W_down"| C[("cache: latent c<br/>d_c numbers")] C -->|"W_uk"| K["keys for every head"] C -->|"W_uv"| V["values for every head"] Qn["query q"] -->|"absorbed: q W_ukᵀ"| QL["query in latent space"] QL -->|"score directly against c"| C

Reading it: the cylinder is all that is stored per token. The two arrows to the right show the plain way: rebuild keys and values from the latent. The bottom path shows the absorbed way: translate the query into latent space once, then score it against the cached latents directly, and mix latents before expanding with W_uv. Both paths give identical outputs; the absorbed one never materializes the big keys and values.

Level 3: the formula and its symbols

$$ c_t = x_t W^{D}, \qquad k_t = c_t W^{UK}, \qquad v_t = c_t W^{UV}, \qquad q \cdot k_t = \big(q\, {W^{UK}}^{\top}\big) \cdot c_t $$

Symbols

Symbol Meaning here Shape / example
$x_t$ token $t$'s vector coming into the layer $d_\text{model}$; (1, 0, 2, 1)
$W^{D}$ the learned "down" projection that squeezes $d_\text{model} \times d_c$
$c_t$ the latent: the only thing cached $d_c$; (3, 1)
$W^{UK}$, $W^{UV}$ learned "up" projections to every head's keys and values $d_c \times H d_h$
$k_t$, $v_t$ token $t$'s keys and values for all heads, side by side $H d_h$; $k$ = (3, 1, 3, 1)
$q$ a query (one head's slice, or all heads side by side) (1, 2, 0, 1)
${W^{UK}}^{\top}$ $W^{UK}$ transposed (rows become columns) $H d_h \times d_c$
$q\,{W^{UK}}^{\top}$ the query translated into latent space $d_c$; (1, 3)

In words: "squeeze each token into a short latent and cache only that; keys and values are the latent times fixed up-projections, so a query can be moved into latent space once and scored against the cache directly."

With the numbers: c = (3, 1), k = (3, 1, 3, 1), q · k = 3 + 2 + 0 + 1 = 6, and (1, 3) · (3, 1) = 3 + 3 = 6.

Level 3: in Python

In Python:

x = [1, 0, 2, 1]
W_down = [[1, 0], [0, 1], [1, 0], [0, 1]]
W_uk = [[1, 0, 1, 0], [0, 1, 0, 1]]
def vecmat(v, M):
    # row vector times matrix: entry c is Σ_r v_r M[r][c]
    return [sum(v[r] * M[r][c] for r in range(len(v))) for c in range(len(M[0]))]
# c = x W_down: the only thing cached
c = vecmat(x, W_down)
c  # → [3, 1]
# k = c W_uk: both heads' keys, rebuilt on demand
k = vecmat(c, W_uk)
k  # → [3, 1, 3, 1]
q = [1, 2, 0, 1]
sum(q_m * k_m for q_m, k_m in zip(q, k))  # → 6
# absorbed: move q into latent space, then score against c
W_uk_T = [list(col) for col in zip(*W_uk)]
q_latent = vecmat(q, W_uk_T)
q_latent, sum(a * b for a, b in zip(q_latent, c))  # → ([1, 3], 6)
# DeepSeek-V2's shape: cached numbers per token per layer
2 * 128 * 128, 512 + 64, round(2 * 128 * 128 / (512 + 64), 1)  # → (32768, 576, 56.9)

Why can so few numbers stand in for so many? Because every head's keys and values are built from the same token, they are highly redundant. In the language of primer.notation, the key projection W^D W^UK has low rank: however many numbers it outputs, they all vary along only d_c independent directions. Storing those d_c coordinates loses nothing that projection can produce. One wrinkle: rotary position embeddings (see primer.ml.positional) rotate each key by its position, which breaks the absorption trick, so DeepSeek-V2 carries position in that separate small 64-number key.

Cache per token on the 32-layer example: 512 KB for 32 KV heads, 128 KB for 8, 36 KB for a 576-number latent, 33 KB for 8 heads at 4 bits, 16 KB for one head

Reading it: each bar is the cache one token costs across all 32 layers of the running example, for one design. The top bar is classic multi-head attention; every bar below it is a way of shrinking it. Sharing heads (GQA, MQA), squeezing into a latent, and cutting bits land in the same range, tens of kilobytes instead of hundreds, by different routes that can be combined.

In code: LatentKVAttention holds the four projections; LatentKVAttention.compress produces the cache, LatentKVAttention.attend expands it into keys and values, and LatentKVAttention.attend_absorbed gets the identical output without expanding. latent_kv_worked_example is the (3, 1) example and latent_kv_bytes_per_token the memory arithmetic.

Quantizing the cache

Everyday picture. The same move as for weights in primer.ml.inference: write each number to the nearest tenth instead of the nearest thousandth, with one ruler per stored vector so one big number does not coarsen all the others.

Tiny worked example. A cached key (0.7, −0.3, 0.2, 0.04) at 4 bits (codes −7 to 7). The step size is 0.7 / 7 = 0.1, the codes are (7, −3, 2, 0), and reading back gives (0.7, −0.3, 0.2, 0): the tiny 0.04 is lost. A 128-number head vector drops from 256 bytes to 64 bytes of codes plus a 2-byte scale: 66 bytes, 3.9 times smaller.

flowchart LR KV["new key or value<br/>16-bit numbers"] -->|"scale = max / 7<br/>code = round(x / scale)"| ST[("cache: 4-bit codes<br/>+ one scale per vector")] ST -->|"code × scale"| R["approximate key or value"] R --> A["attention as usual"]

Reading it: writing to the cache quantizes once per token; reading from it multiplies back. Attention itself is unchanged; it just sees slightly rounded keys and values. The saving is on the cylinder, the part that grows with every token.

Level 3: the formula and its symbols

$$ \text{bytes per token} = 2\,L\,H_{kv}\left(d_h \cdot \frac{\text{bits}}{8} + \frac{\text{scale bits}}{8}\right) $$

Symbols

Symbol Meaning here In the example
$2\,L\,H_{kv}$ how many head vectors one token stores: a key and a value per KV head per layer 2 × 32 × 8 = 512
$d_h$ numbers per head vector 128
bits / 8 bytes per code 4-bit: 0.5
scale bits / 8 bytes for the one scale each vector carries 16-bit: 2

In words: "every stored head vector costs its codes plus one scale, and a token stores a key vector and a value vector per head per layer."

With the numbers: 2 × 32 × 8 × (128 × 0.5 + 2) = 512 × 66 = 33,792 bytes, against 131,072 at 16 bits: 3.9 times smaller.

Level 3: in Python

In Python:

k = [0.7, -0.3, 0.2, 0.04]
# one scale per vector: the largest magnitude lands on code 7
s = max(abs(k_j) for k_j in k) / 7
round(s, 3)  # → 0.1
codes = [round(k_j / s) for k_j in k]
codes  # → [7, -3, 2, 0]
[round(s * c, 2) for c in codes]  # → [0.7, -0.3, 0.2, 0.0]
L, H_kv, d_h = 32, 8, 128
# bytes per token: 16-bit, then 4-bit codes plus a 16-bit scale per vector
2 * L * H_kv * d_h * 16 / 8, 2 * L * H_kv * (d_h * 4 / 8 + 16 / 8)  # → (131072.0, 33792.0)

Why it matters in practice. On random test data, an 8-bit cache moves the attention output by under 1% and a 4-bit cache by about 12%; trained models tolerate this far better than random data suggests, and careful schemes go lower. Keys tend to have a few channels with consistently large values, so KIVI quantizes keys per channel and values per token and reaches 2 bits. Quantization stacks with everything above: a GQA model with a sliding window and a 4-bit cache enjoys all three savings.

In code: quantize_kv quantizes each cached vector with its own scale via primer.ml.inference.quantize, quantized_kv_bytes_per_token is the formula, and quantized_cache_error measures how far attention's output moves.

Putting it together

On log axes from 1k to 1M tokens, multi-head and grouped-query caches climb past 80 GB, latent and 4-bit caches climb 4 times lower, the hybrid 8 times lower, while the sliding window flattens at 0.5 GB and the SSM state stays at 4 MB

Reading it: the x-axis is context length and the y-axis is cache memory for one conversation, both logarithmic, with the dashed line at 80 GB. Lines with the same slope grow the same way (in proportion to n); compression moves a line down without changing its slope. The sliding window bends flat at 4,096 tokens, and the SSM is flat from the start. Flat lines are what make million-token contexts affordable; the price is that they no longer keep every token exactly.

Technique What it cuts What it gives up
Sliding window compute to n·w, cache to w tokens direct access beyond the window
Sparse (global, strided) compute exact long-range pairs not in the pattern
Linear attention compute to n·d², cache to a fixed state sharp, exact recall
SSM / Mamba the same the same, softened by selectivity
Hybrid most of the cache a little of both
GQA / MQA cache by the sharing factor a little modelling capacity
Latent KV cache by d_c / (2·H·d_h) extra matrix work, care with positions
Quantized KV cache by 16 / bits a little precision

In code: cache_bytes_by_method computes every line in the figure for a given context length.

In 20 seconds

  • Long context has two bills: attention scores grow with n², and the KV cache grows with n (per layer, per conversation). At 1M tokens the cache of an 8B-class model alone is about 137 GB.
  • Sliding window: each token reads the last w; stacked layers still reach L·(w − 1) back, and a rolling buffer caps the cache at w tokens. Sparse patterns add a few global or strided links so any two tokens are a hop or two apart.
  • Linear attention replaces e^(q·k) with φ(q)·φ(k), which turns attention into a running sum: O(n) work and a fixed-size state, but blurrier recall.
  • State-space models update a fixed state linearly; fixed ones train as a convolution, selective ones (Mamba) let each token set how much to keep and write, and train with a parallel scan. Hybrids keep a few attention layers for exact lookup.
  • Compress the cache: share KV heads (GQA/MQA), cache a small latent and expand on demand (latent KV), and store fewer bits per number.

Self-test questions

Why does doubling the context quadruple attention's compute but only double its cache? Every token's query is scored against every key, so the scores form an n × n table: doubling n quadruples it. The cache stores one key and one value per token per layer, a list that grows by one entry per token, so doubling n doubles it.

A model uses a 4,096-token sliding window in all 32 layers. Can token 100,000 be influenced by token 1? Yes, in principle: information moves up to w − 1 = 4,095 positions per layer, so 32 layers reach 131,040 positions back. In practice it must be relayed through about 25 intermediate tokens and layers, so it arrives weakened; direct, exact lookup only works within the window.

What does a global token do in a sparse pattern, and why is it cheap? Every token reads it and it reads every token, so any two tokens are at most two hops apart. It adds only about one column and one row of scores, a cost that grows with n rather than n².

Why can linear attention run as a recurrence but softmax attention cannot? Linear attention's weight φ(q)·φ(k) splits into a query part and a key part, so the key-and-value parts can be summed ahead of time into a fixed-size state that any later query can read. The softmax weight e^(q·k) does not split that way, so each new query must revisit every stored key.

What does "selective" mean in Mamba, and what problem does it fix? The step size Δ, and with it how much of the old state is kept and how much of the new token is written, is computed from each token. A fixed SSM applies the same keep and write amounts to every token, so it cannot both absorb one important token and ignore the filler around it.

If a selective SSM's parameters change every step, how does it train in parallel? The update "multiply by a, add b" can be merged: two consecutive steps form one step of the same kind. A parallel scan merges pairs, then pairs of pairs, and finishes in about log₂ n rounds.

How does latent KV caching save memory without changing the attention outputs? Keys and values are computed as a small cached latent times fixed up-projection matrices, so storing the latent is enough to rebuild them exactly. The up-projection can even be folded into the query and output side, so the full keys and values are never built.

Why do hybrid models keep a few attention layers instead of going all-SSM? A fixed-size state is a lossy summary, which hurts copying and exact recall over long contexts. A handful of attention layers restores exact lookup while most layers keep constant memory, so the cache shrinks roughly by the fraction of layers that are SSMs.

The papers behind this lesson

  • Child, Gray, Radford & Sutskever, Generating Long Sequences with Sparse Transformers (2019): https://arxiv.org/abs/1904.10509. Introduced strided and fixed sparse attention patterns, cutting attention's cost to about n√n.
  • Beltagy, Peters & Cohan, Longformer: The Long-Document Transformer (2020): https://arxiv.org/abs/2004.05150. Combined a sliding window with task-chosen global tokens to read long documents in linear time.
  • Zaheer et al., Big Bird: Transformers for Longer Sequences (2020): https://arxiv.org/abs/2007.14062. Mixed window, global and random links, and proved such sparse attention keeps the expressive power of full attention.
  • Katharopoulos et al., Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention (2020): https://arxiv.org/abs/2006.16236. Replaced softmax with a kernel feature map, turning attention into a running sum with O(n) cost.
  • Gu, Goel & Ré, Efficiently Modeling Long Sequences with Structured State Spaces (S4, 2021): https://arxiv.org/abs/2111.00396. Made linear state-space layers trainable on very long sequences through the convolution view.
  • Gu & Dao, Mamba: Linear-Time Sequence Modeling with Selective State Spaces (2023): https://arxiv.org/abs/2312.00752. Made the state-space parameters depend on the input and trained them with a hardware-aware parallel scan. Annotated companion
  • Dao & Gu, Transformers are SSMs (Mamba-2, 2024): https://arxiv.org/abs/2405.21060. Showed that selective SSMs and a form of linear attention are two views of one computation.
  • Jiang et al., Mistral 7B (2023): https://arxiv.org/abs/2310.06825. Used sliding-window attention with a rolling-buffer cache in a strong open model. Annotated companion
  • Xiao et al., Efficient Streaming Language Models with Attention Sinks (2023): https://arxiv.org/abs/2309.17453. Found that models lean on the first few tokens, and that keeping them plus a window allows streaming without limit.
  • Lieber et al., Jamba: A Hybrid Transformer-Mamba Language Model (2024): https://arxiv.org/abs/2403.19887. Interleaved one attention layer per seven Mamba layers to cut the KV cache while keeping recall.
  • Shazeer, Fast Transformer Decoding: One Write-Head is All You Need (2019): https://arxiv.org/abs/1911.02150. Introduced multi-query attention, one shared key/value head for all query heads. Annotated companion
  • DeepSeek-AI, DeepSeek-V2 (2024): https://arxiv.org/abs/2405.04434. Introduced multi-head latent attention, caching a small latent per token instead of full keys and values.
  • Liu et al., KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache (2024): https://arxiv.org/abs/2402.02750. Quantized keys per channel and values per token, bringing the cache down to 2 bits.

Further reading

on GitHub
   1r"""
   2# Long context and efficient architectures: reading a million tokens without a million-squared bill
   3
   4Run: `python -m primer.ml.efficient_architectures`
   5
   6New to the notation? `primer.notation` explains every symbol used here from
   7zero. This lesson builds on attention from `primer.ml.attention` and the KV
   8cache from `primer.ml.inference`.
   9
  10## Level 1: The practitioner's guide
  11
  12**In one sentence.** Long context is paid for twice, in compute that grows
  13with the square of the length and in cache memory that grows with the
  14length, and every "efficient" architecture is a different bargain about
  15which tokens a model keeps exactly, which it summarizes, and how compactly
  16it stores them.
  17
  18**When you need it.** You need this when a model's context window is part
  19of your design: a whole codebase in the prompt, a day-long conversation, a
  20document set that will not fit in retrieval, or an agent that runs for
  21hundreds of steps. You also need it when choosing between models whose
  22cards list different attention designs (grouped-query, sliding-window,
  23hybrid Mamba, latent attention), because those words decide what the model
  24will cost you per conversation and what it will forget. The tell is a
  25context length that sounds like a solved problem. On this lesson's
  26Llama-3-8B-shaped model, a 128,000-token conversation holds 17.2 GB of
  27cache and a million-token one 137 GB, more than an 80 GB GPU before any
  28weights are loaded, and the attention work grows 16,384-fold across that
  29same range (this lesson's context-cost table). You do not need this lesson
  30for prompts of a few thousand tokens: at that length full attention is
  31cheap and exact, and nothing here beats it.
  32
  33**Your options.** Each is a design you choose by choosing a model, or a
  34setting you apply to one, from the most exact to the most compact:
  35
  36| Option | What it does | What it keeps | What it costs | Where it lives |
  37|---|---|---|---|---|
  38| Full attention with grouped-query heads | Every token reads every earlier token; several query heads share each key/value head | Exact recall of everything in context | 128 KB per token with 8 KV heads on the 8B example, four times less than 32 heads; still 17 GB at 128k | The model's architecture |
  39| Sliding-window and interleaved layers | Each token reads the last w tokens; depth relays information further; some models keep a few full-attention layers | Exact recall inside the window, relayed and weaker beyond it | A cache capped at w tokens: 0.54 GB for a 4,096 window on the 8B example, whatever the length | The model's architecture |
  40| Attention sinks with a window | Keeps the first few tokens plus a recent window at serving time | Fluent streaming without limit; no recall of the dropped middle | Nearly nothing; StreamingLLM reports streaming to 4 million tokens and up to 22.2× speedup over recomputation | The inference server |
  41| Hybrid state-space and attention | Most layers carry a fixed-size state; one in several is attention | Exact lookup through the attention layers, cheap everything else | On the 8B shape, 4 attention layers in 32 cut the 128k cache from 17.2 GB to 2.15 GB (this lesson's hybrid formula) | The model's architecture |
  42| Pure state-space or linear attention | Every layer summarizes the past into a fixed state | Fluent long text; blurrier exact recall | A cache that never grows (4 MB for the whole 8B-shaped stack in this lesson's demo); weaker copying of exact tokens from far back | The model's architecture |
  43| Latent KV cache | Caches a short latent per token and rebuilds keys and values from it | Exact outputs, by construction | DeepSeek-V2 caches 576 numbers per token per layer instead of 32,768 (about 57× less) at the price of extra matrix work | The model's architecture |
  44| A quantized KV cache | Stores cached keys and values at 8 or 4 bits with one scale per vector | Slightly rounded attention | 3.9× smaller at 4 bits; under 1% output error at 8 bits and about 12% at 4 bits on random data (this lesson's measurement) | The inference server |
  45
  46**How to choose.** Start from what the task needs to find, not from the
  47advertised window.
  48
  49- Exact lookups across a long input (a function name in a repository, a
  50  clause in a contract): keep full attention in enough layers. A
  51  grouped-query model, or a hybrid with attention layers, and a budget for
  52  the cache.
  53- Long, fluent, forward-moving text (a running transcript, a stream) with
  54  no need to quote the distant past: a sliding-window or state-space model
  55  is far cheaper, and a sink-plus-window server setting keeps even a
  56  full-attention model streaming.
  57- Many concurrent long conversations on fixed hardware: the cache per
  58  token is your capacity. Prefer fewer KV heads or a latent cache, then
  59  quantize the cache, then cap the context you allow.
  60- A context window claimed at 128k or more: test the model at the length
  61  you will use. RULER (Hsieh et al., 2024) found that of 17 models claiming
  62  32k tokens or more, only half held up at 32k, and *Lost in the Middle*
  63  (Liu et al., 2023) found accuracy highest when the relevant passage sits
  64  at the start or the end of the input and worst in the middle.
  65- Whatever you pick, put the facts the model must use where the design
  66  keeps them exactly: inside the window, near the ends, or in the prompt of
  67  a retrieval step (`primer.agents.rag`) instead of a million-token dump.
  68
  69**What it costs.** Two meters run at once. Compute is the square of the
  70length and is paid at prefill: on 4,096 tokens, full attention scores 32
  71times as many pairs as a 64-token window (this lesson's demo). Memory is
  72linear in the length and is paid for as long as a conversation stays open;
  73it decides how many users a GPU serves. The compressions differ in what
  74they charge: grouped-query heads cost a little modelling capacity; the
  75window costs direct access beyond it; a fixed state costs sharp recall
  76(in this lesson's comparison, softmax attention puts up to 0.46 of a row
  77on its favourite key where linear attention manages 0.21); the latent
  78cache costs extra matrix work and care with positions; quantization costs
  79precision. The gains are large: Mamba reports 5× the inference throughput
  80of a transformer with linear scaling in length, DeepSeek-V2 reports a KV
  81cache 93.3% smaller and 5.76× the generation throughput of its
  82predecessor, and Jamba fits a 256k-context model on one 80 GB GPU (each
  83paper's abstract).
  84
  85**What breaks.**
  86
  87- **The fact was outside the window.** A sliding-window model can be
  88  influenced by a token 131,040 positions back (Mistral 7B's 32 layers of
  89  4,096), but only by relay through about 25 hops, so a distant fact
  90  arrives weakened. Do not expect exact quotes from beyond the window.
  91- **A plain window drops the first tokens and the model falls apart.**
  92  Trained models park attention on the first few tokens (attention sinks);
  93  keep them when you truncate.
  94- **A fixed-size state cannot copy.** Pure state-space and linear-attention
  95  models do worse at reproducing a specific token from far back. If the
  96  task is retrieval or copying, keep attention layers.
  97- **The advertised context is not the effective context.** Test at your
  98  length with your task, in the middle of the input, not only with a
  99  needle at the end.
 100- **A sparse pattern that saves no time.** Skipped pairs only save work
 101  when the kernel skips whole blocks of the score matrix; a pattern drawn
 102  token by token costs as much as full attention.
 103- **A quantized cache that changes answers.** Keys carry a few channels
 104  with large values; quantizing per token can lose them. Check outputs on
 105  your own data, and use a scheme like KIVI's (keys per channel, values
 106  per token) at low bit widths.
 107
 108**In the wild.** Mistral 7B (Jiang et al.) shipped sliding-window attention
 109with a rolling-buffer cache; Llama 3 uses grouped-query attention, the
 110design this lesson's 8 KV heads mirror (Grattafiori et al., in
 111`primer.ml.inference`). Jamba (Lieber et al.) interleaves one attention
 112layer among Mamba layers; Mamba itself (Gu and Dao) and Mamba-2 (Dao and
 113Gu) are the state-space models; the Sparse Transformer (Child et al.),
 114Longformer (Beltagy et al.) and BigBird (Zaheer et al.) are the sparse
 115patterns; StreamingLLM (Xiao et al.) is the sink-plus-window trick;
 116DeepSeek-V2 introduced the latent cache; KIVI (Liu et al.) reaches a 2-bit
 117cache. RULER and *Lost in the Middle* are the two tests to run before
 118trusting a context length. Inference servers expose the serving-side
 119options: vLLM and SGLang list quantization and prefix caching among their
 120features, and the paged KV memory both use is what makes a growing cache
 121manageable at all (`primer.ml.inference`).
 122
 123**Go deeper.** Level 2 builds every row of the table from nothing: the two
 124bills counted for real context lengths, the window mask and the reach it
 125gives a stack of layers, sparse patterns with their pair counts, linear
 126attention as a running sum, a state-space model that is both a loop and a
 127convolution, Mamba's per-token step size on a recall task, the parallel
 128scan that trains it, the latent cache with its absorbed query, and a
 129quantized cache with its error measured. If you only needed to pick a
 130model and size its cache, you are done.
 131
 132## Level 2: How it works, from scratch
 133
 134Attention (see `primer.ml.attention`) lets every token look at every earlier
 135token. That is where its power comes from, and it is also where its bill comes
 136from. The bill has two lines:
 137
 138- **Compute.** Every token is scored against every other token, so the work
 139  grows with the *square* of the context length n.
 140- **Memory.** During generation the model keeps every earlier token's keys
 141  and values (the KV cache, see `primer.ml.inference`), so memory grows with
 142  n, per layer, per conversation.
 143
 144Everything in this lesson attacks one of those two lines, in one of three ways:
 145
 1461. **Look at fewer tokens**: sliding-window and sparse attention.
 1472. **Summarize the past into a fixed-size state**: linear attention and
 148   state-space models (Mamba), which run like a recurrent network.
 1493. **Store the past more compactly**: fewer key/value heads, a small latent
 150   vector per token, fewer bits per number.
 151
 152```mermaid
 153flowchart TB
 154  P["Long context is expensive<br/>compute grows with n², memory with n"] --> F["Look at fewer tokens"]
 155  P --> S["Keep a fixed-size summary"]
 156  P --> C["Store the past compactly"]
 157  F --> F1["sliding window"] & F2["sparse: local + global, strided"]
 158  S --> S1["linear attention"] & S2["state-space models, Mamba"]
 159  C --> C1["GQA / MQA"] & C2["latent KV"] & C3["quantized KV"]
 160  S2 --> H["hybrids: a few attention layers<br/>among many SSM layers"]
 161  F1 --> H
 162```
 163
 164**Reading it:** start at the top box, the problem. The three arrows below it
 165are the three strategies, and the leaves are the techniques this lesson
 166builds from scratch, section by section. The bottom box is where real
 167models often land: a mix, keeping a little full attention for exact lookups
 168and using cheap layers for everything else.
 169
 170## 1. Why long context is expensive
 171
 172**Everyday picture.** A dinner party where every new guest must shake hands
 173with every guest already there. Ten guests make 45 handshakes; a thousand
 174guests make half a million. On top of that, the coat check keeps one coat per
 175guest on every floor of the building, so the coat racks grow with every
 176arrival. Attention pays both bills: the handshakes are the query-key scores,
 177and the coats are the KV cache.
 178
 179**Tiny worked example.** Take a Llama-3-8B-shaped model: 32 layers, 8
 180key/value heads, 128 numbers per head, 16-bit numbers (2 bytes). One token
 181costs 2 × 32 × 8 × 128 × 2 = 131,072 bytes of cache, about 128 KB.
 182
 183| Context n | Pairs scored per head per layer (n²) | KV cache for one conversation |
 184|---|---|---|
 185| 8,192 (8k) | 67 million | 1.1 GB |
 186| 131,072 (128k) | 17 billion | 17.2 GB |
 187| 1,048,576 (1M) | 1.1 trillion | 137.4 GB |
 188
 189Going from 8k to 1M is 128 times more tokens, 16,384 times more pairs, and a
 190cache that no longer fits on one 80 GB GPU.
 191
 192```mermaid
 193flowchart LR
 194  T["new token n"] --> Q["its query"]
 195  Q -->|"scored against"| KC[("KV cache<br/>n − 1 keys and values<br/>per layer")]
 196  KC --> O["output for token n"]
 197  T -->|"its key and value<br/>are appended"| KC
 198```
 199
 200**Reading it:** follow one new token. Its query is scored against every key
 201already in the cache, so the work for this one token grows with n, and the
 202work for all n tokens grows with n². Then its own key and value are appended,
 203so the cache grows by one entry, in every layer, for every token. The two
 204arrows into the cache are the two bills.
 205
 206$$
 207\text{pairs} = n^2 \qquad\qquad \text{KV bytes} = 2 \, L \, H_{kv} \, d_h \, b \, n
 208$$
 209
 210**Symbols**
 211
 212| Symbol | Meaning here | In the example |
 213|---|---|---|
 214| $n$ | context length: how many tokens are in play | 131,072 |
 215| $n^2$ | $n$ times $n$: every query paired with every key (causal masking halves it, which does not change how it grows) | 17,179,869,184 |
 216| $2$ | one key and one value per token | |
 217| $L$ | layers, each with its own cache | 32 |
 218| $H_{kv}$ | key/value heads per layer | 8 |
 219| $d_h$ | numbers per head | 128 |
 220| $b$ | bytes per stored number | 2 (16-bit) |
 221
 222**In words:** "the score work is the context length squared, and the cache
 223holds a key and a value for every head, in every layer, for every token."
 224
 225**With the numbers:** at n = 131,072 the pairs are 131,072² ≈ 1.7 × 10¹⁰,
 226and the cache is 2 × 32 × 8 × 128 × 2 × 131,072 = 17,179,869,184 bytes ≈
 22717.2 GB.
 228
 229**In Python:**
 230
 231```python
 232L, H_kv, d_h, b = 32, 8, 128, 2
 233# bytes per token: a key and a value, per layer, per KV head
 234per_token = 2 * L * H_kv * d_h * b
 235per_token  # → 131072
 236contexts = [8_192, 131_072, 1_048_576]
 237# pairs per head per layer
 238[f"{n * n:,}" for n in contexts]  # → ['67,108,864', '17,179,869,184', '1,099,511,627,776']
 239# GB of cache for one conversation
 240[round(n * per_token / 1e9, 1) for n in contexts]  # → [1.1, 17.2, 137.4]
 241```
 242
 243![Pairs scored grow from 67 million at 8k tokens to 1.1 trillion at 1M, and the KV cache from 1.1 GB to 137 GB, past an 80 GB GPU](figures/primer.ml.efficient_architectures.context_cost.svg)
 244
 245**Reading it:** the left panel counts query-key pairs per head per layer, on
 246a logarithmic axis where each gridline is ten times the last: every step to
 247the right multiplies the bar by 256 and then by 64. The right panel is the KV
 248cache for one conversation on an ordinary axis, with the dashed line at 80 GB,
 249the memory of one large GPU. The 1M bar crosses it before any weights are
 250loaded. The left panel is the compute bill; the right panel is the memory
 251bill.
 252
 253**Why it matters in practice.** The compute bill is paid once per prompt
 254(prefill), and exact tricks such as FlashAttention make it faster without
 255changing it. The memory bill is paid for as long as a conversation is open,
 256and it decides how many users one GPU can serve. That is why most of the
 257techniques below target the cache.
 258
 259**In code:** `context_cost` returns both lines of the bill for any context length, using `primer.ml.inference.kv_cache_bytes` for the cache.
 260
 261## 2. Sliding-window attention: reading through a letterbox
 262
 263**Everyday picture.** You read a long scroll through a letterbox slot that
 264shows only the last w words. On your own you would lose the beginning. But
 265suppose a row of readers sits one above the other, each reading the notes of
 266the reader below through a slot of the same width. Each reader passes things
 267along a little further, like a bucket brigade, so a stack of readers can
 268carry a fact much further back than any one slot shows.
 269
 270**Tiny worked example.** Eight tokens, window w = 3: each token reads itself
 271and the two tokens before it.
 272
 273```text
 274            key:  0 1 2 3 4 5 6 7
 275query 0           ✓ · · · · · · ·
 276query 1           ✓ ✓ · · · · · ·
 277query 2           ✓ ✓ ✓ · · · · ·
 278query 3           · ✓ ✓ ✓ · · · ·
 279query 4           · · ✓ ✓ ✓ · · ·
 280query 5           · · · ✓ ✓ ✓ · ·
 281query 6           · · · · ✓ ✓ ✓ ·
 282query 7           · · · · · ✓ ✓ ✓
 283```
 284
 285That is 1 + 2 + 3 × 6 = **21** scores instead of the 36 a full causal mask
 286needs. Now stack three such layers. Token 7 reads token 5 in the top layer;
 287token 5 had read token 3 in the layer below; token 3 had read token 1 in the
 288layer below that. So after three layers, token 1 can influence token 7, six
 289positions back, but token 0 cannot.
 290
 291```mermaid
 292flowchart RL
 293  subgraph L3["layer 3"]
 294    a7["token 7"]
 295  end
 296  subgraph L2["layer 2"]
 297    b5["token 5"]
 298  end
 299  subgraph L1["layer 1"]
 300    c3["token 3"]
 301  end
 302  subgraph IN["input"]
 303    d1["token 1"]
 304    d0["token 0: out of reach"]
 305  end
 306  a7 -->|"reads 2 back"| b5 -->|"reads 2 back"| c3 -->|"reads 2 back"| d1
 307```
 308
 309**Reading it:** read right to left, from token 7 at the top layer down to the
 310input. Each arrow is one layer's window: it can reach at most w − 1 = 2
 311positions back. Three layers chain three arrows, so the furthest input token 7
 312can hear from is 3 × 2 = 6 positions back, token 1. Token 0 would need a
 313fourth hop. Information still travels far, but only in stages.
 314
 315$$
 316M_{ij} = \begin{cases} 1 & \text{if } 0 \le i - j < w \\ 0 & \text{otherwise} \end{cases}
 317\qquad\qquad
 318\text{reach} = L\,(w - 1)
 319$$
 320
 321**Symbols**
 322
 323| Symbol | Meaning here | In the example |
 324|---|---|---|
 325| $M_{ij}$ | the mask: 1 if query $i$ may read key $j$, 0 if not | row 5: keys 3, 4, 5 |
 326| $i$, $j$ | the query's position and the key's position | $i = 5$ |
 327| $i - j$ | how far back the key is; negative means the future | 0, 1, 2 allowed |
 328| $w$ | the window: how many tokens each query reads, itself included | 3 |
 329| $\{$ | "cases": use the line whose condition holds | |
 330| $L$ | how many windowed layers are stacked | 3 |
 331| reach | the furthest back an input can influence an output | 6 |
 332
 333**In words:** "a query may read a key if the key is not in the future and is
 334fewer than w positions back; stacking L layers lets information travel
 335L times w − 1 positions."
 336
 337**With the numbers:** row 5 allows j = 3, 4, 5 (5 − 3 = 2 < 3). Three layers
 338of w = 3 reach 3 × 2 = 6. Mistral 7B uses w = 4,096 over 32 layers:
 33932 × 4,095 = 131,040 tokens, about 131k, from a window of 4k.
 340
 341**In Python:**
 342
 343```python
 344n, w = 8, 3
 345# M_ij for the row of token 5
 346[int(0 <= 5 - j < w) for j in range(n)]  # → [0, 0, 0, 1, 1, 1, 0, 0]
 347# pairs scored: row i keeps min(i + 1, w) keys
 348sum(min(i + 1, w) for i in range(n))  # → 21
 349# full causal attention, for comparison
 350n * (n + 1) // 2  # → 36
 351# reach after L layers
 352L = 3
 353L * (w - 1)  # → 6
 354# Mistral 7B: 32 layers, a 4,096-token window
 35532 * (4096 - 1)  # → 131040
 356```
 357
 358![Left, one layer with a window of 4 is a thin diagonal band; right, after three layers each token can hear the last 10 positions, a band three times as wide](figures/primer.ml.efficient_architectures.window_reach.svg)
 359
 360**Reading it:** both panels are 16 × 16 grids, with rows for the token doing
 361the looking and columns for the token being looked at, like the causal
 362heatmap in `primer.ml.attention`. Dark means "can reach". On the left, one
 363layer with w = 4 is a thin band along the diagonal: 4 cells per row at most.
 364On the right, the same mask applied three times: the band has widened to
 3653 × 3 + 1 = 10 cells, because each layer extends the reach by w − 1 = 3.
 366Neither panel ever touches the upper-right triangle: the window is still
 367causal.
 368
 369**In code:** `sliding_window_mask` builds the band, `pairs_computed` counts its cells, `receptive_field` applies the mask layer after layer to find who can reach whom, and `reach` is the L(w − 1) formula.
 370
 371### The rolling buffer: a cache that stops growing
 372
 373**Everyday picture.** A whiteboard with room for exactly w notes. When it is
 374full, the next note goes over the oldest one. You never need a bigger board,
 375however long the meeting runs.
 376
 377**Tiny worked example.** With w = 4, after token 9 the buffer holds the keys
 378and values of tokens 6, 7, 8 and 9. Token 10 overwrites token 6's slot. On
 379the running 8B example at 128k tokens, a 4,096-token window holds 4,096 × 128
 380KB = 0.54 GB instead of 17.2 GB: **32 times less**, and the same at 1M tokens.
 381
 382```mermaid
 383flowchart LR
 384  subgraph B["buffer of w = 4 slots, after token 9"]
 385    s0["slot 0: token 8"]
 386    s1["slot 1: token 9"]
 387    s2["slot 2: token 6 (oldest)"]
 388    s3["slot 3: token 7"]
 389  end
 390  N["token 10 arrives"] -->|"overwrites the oldest"| s2
 391  B --> A["token 10 attends to<br/>tokens 7, 8, 9, 10"]
 392```
 393
 394**Reading it:** the four slots are all the memory there is. Slot number is
 395token number modulo 4 (the remainder after dividing by 4), so token 10 lands
 396in slot 2, where token 6 sat. Token 6 was about to leave the window anyway,
 397so nothing the window needs is lost.
 398
 399**Why it matters in practice.** Mistral 7B combined the window with this
 400rolling buffer. Several later model families interleave sliding-window layers
 401with a few full-attention layers: the local layers are cheap, and the full
 402layers keep long-range lookups exact. The weakness is plain from the diagram:
 403a fact beyond the reach is invisible, and a fact inside the reach has to
 404survive several hops.
 405
 406**In code:** `sliding_window_decode` generates token by token with a buffer that never holds more than w keys, and returns the same outputs as masked attention with `sliding_window_mask`.
 407
 408## 3. Sparse attention: a few long-distance lines
 409
 410**Everyday picture.** An open-plan office. You mostly talk to the people at
 411the desks next to yours (local). Anyone can phone the front desk, and the
 412front desk hears from everyone (a *global* token): any two people are at
 413most two calls apart. Another design is the express train: you talk to your
 414neighbours and also to every fourth desk down the row (*strided*), so a
 415message can travel far in a few big jumps.
 416
 417**Tiny worked example.** Sixteen tokens, window 4.
 418
 419- The window alone: 1 + 2 + 3 + 4 × 13 = **58** pairs.
 420- Make token 0 global: the 12 rows past the window (tokens 4 to 15) add a
 421  pair each: **70** pairs, against **136** for full causal attention.
 422- Strided with stride 4: the last 4 tokens, plus every 4th token before them.
 423  Token 13 reads 10, 11, 12, 13 and then 9, 5, 1. Over all rows that is
 424  58 + 24 = **82** pairs.
 425
 426```mermaid
 427flowchart LR
 428  G(("token 0<br/>global"))
 429  t3["token 3"] --- G
 430  t7["token 7"] --- G
 431  t11["token 11"] --- G
 432  t15["token 15"] --- G
 433  t14["token 14"] --- t15
 434  t13["token 13"] --- t14
 435```
 436
 437**Reading it:** the circle is the global token. Every token has a line to
 438it, so token 15 can reach token 3 in two hops, through token 0, however long
 439the sequence. The short lines at the bottom are the ordinary local window.
 440The picture is sparse (few lines), yet no two tokens are far apart.
 441
 442$$
 443M_{ij} = 1 \quad\text{when}\quad i - j \ge 0 \;\text{ and }\; \big(\, i - j < s \;\text{ or }\; (i - j) \bmod s = 0 \,\big)
 444$$
 445
 446**Symbols**
 447
 448| Symbol | Meaning here | In the example |
 449|---|---|---|
 450| $M_{ij}$ | 1 if query $i$ may read key $j$ | |
 451| $i - j$ | how far back key $j$ is | for $i = 13$: 0 to 13 |
 452| $s$ | the stride: the local width, and the jump between long-range keys | 4 |
 453| $\bmod$ | "modulo": the remainder after dividing; $(i - j) \bmod s = 0$ means "a whole number of strides back" | 8 mod 4 = 0 |
 454| and, or | both conditions must hold; at least one must hold | |
 455
 456**In words:** "never read the future; read the last s tokens, and beyond
 457them every s-th token."
 458
 459**With the numbers:** for i = 13, s = 4: distances 0 to 3 give keys 13, 12,
 46011, 10; distances 4, 8 and 12 give keys 9, 5, 1.
 461
 462**In Python:**
 463
 464```python
 465s = 4
 466# row 13: which keys does token 13 read?
 467[j for j in range(16) if 13 - j >= 0 and (13 - j < s or (13 - j) % s == 0)]  # → [1, 5, 9, 10, 11, 12, 13]
 468# strided pairs over all 16 rows
 469sum(1 for i in range(16) for j in range(i + 1) if i - j < s or (i - j) % s == 0)  # → 82
 470# a window of 4 plus global token 0
 471sum(1 for i in range(16) for j in range(i + 1) if i - j < 4 or j == 0)  # → 70
 472```
 473
 474![Four 16 by 16 masks: full causal attention scores 136 pairs, the window 58, window plus a global first token 70, strided 82](figures/primer.ml.efficient_architectures.sparse_masks.svg)
 475
 476**Reading it:** each panel is a mask, rows for queries and columns for keys,
 477dark where a score is computed; the title gives the count. The first panel
 478is full causal attention, a solid triangle. The window is a diagonal band.
 479Adding a global token fills in the first column: everyone reads token 0. The
 480strided pattern adds dotted diagonals every 4 columns: the express stops. At
 48116 tokens the savings look modest; with a fixed window they grow with n,
 482because the band stays w wide while the triangle keeps growing.
 483
 484**Why it matters in practice.** The Sparse Transformer introduced strided
 485patterns; Longformer and BigBird combined windows with global tokens (BigBird
 486added a few random links too) to read documents of thousands of tokens.
 487StreamingLLM found that trained models park a lot of attention on the very
 488first tokens (*attention sinks*); a plain window drops them and falls apart,
 489while keeping the first few tokens plus a window lets a model stream
 490indefinitely. One caution: a sparse pattern only saves time if the GPU
 491kernel skips whole blocks of the score matrix, which is why real patterns
 492are built from blocks.
 493
 494**In code:** `global_local_mask` adds global rows and columns to a window, `strided_mask` builds the express-stop pattern, and `pairs_computed` counts what each one scores. Any of them can be passed as the mask to `primer.ml.attention.scaled_dot_product_attention`, and every skipped pair gets a weight of exactly 0.
 495
 496## 4. Linear attention: a pot instead of a guest list
 497
 498**Everyday picture.** A potluck soup. In softmax attention every new guest
 499tastes every dish on the table, one by one, and then mixes a bowl: the more
 500dishes, the longer it takes. In linear attention every guest pours their dish
 501into one shared pot as they arrive, and a new guest takes a single ladle,
 502seasoned to their own taste. The pot never grows, and the ladle costs the
 503same for guest 3 as for guest 3 million.
 504
 505The trick that makes the pot possible: softmax's score $e^{q \cdot k}$ cannot
 506be split into "a part that depends on q" times "a part that depends on k".
 507Replace it with $\phi(q) \cdot \phi(k)$, a dot product of transformed
 508vectors, and it can. Then the order of the matrix multiplies can be swapped:
 509$(\phi(Q)\phi(K)^\top)\,V = \phi(Q)\,(\phi(K)^\top V)$. The left side builds an
 510n × n matrix; the right side builds a small d × d one (the pot) and never
 511builds the big one. Matrix multiplication allows regrouping like this
 512(*associativity*), which `primer.notation` covers under matrix multiply.
 513
 514**Tiny worked example.** Three tokens with 2-number keys (1, 0), (0, 1),
 515(1, 1) and one-number values 2, 4, 6. The third token's query is (1, 0). Use
 516the feature map φ(x) = elu(x) + 1, which for a positive number is simply
 517x + 1 and for zero or a negative number is $e^x$ (always above zero).
 518
 5191. φ of the keys: (2, 1), (1, 2), (2, 2).
 5202. Pour into the pot: S = (2, 1)·2 + (1, 2)·4 + (2, 2)·6 = (20, 22), and the
 521   running total of keys z = (5, 5).
 5223. Ladle with φ(q) = (2, 1): (2·20 + 1·22) / (2·5 + 1·5) = 62 / 15 = **4.13**.
 523
 524The slow way agrees: the weights φ(q)·φ(k) are 5, 4 and 6, so the output is
 525(5·2 + 4·6 + 6·6) / 15 = 62/15. Softmax attention on the same numbers gives
 526**4.00**: linear attention is a *different* attention, not a faster copy of
 527the same one.
 528
 529```mermaid
 530flowchart LR
 531  subgraph T["each new token i"]
 532    K["φ(k_i)"]
 533    V["v_i"]
 534    Qi["φ(q_i)"]
 535  end
 536  K & V -->|"add φ(k_i) v_iᵀ"| S[("pot S<br/>d_k × d_v numbers")]
 537  K -->|"add φ(k_i)"| Z[("total z<br/>d_k numbers")]
 538  Qi --> R["output = φ(q_i)ᵀ S / φ(q_i)ᵀ z"]
 539  S --> R
 540  Z --> R
 541```
 542
 543**Reading it:** each token does two things. Its key and value go into the
 544pot S (and its key into the running total z, used to normalize). Its query
 545reads the pot once. S and z have a fixed size set by the head width, not by
 546the number of tokens, so this is a recurrent network: a state updated once
 547per token, with no cache that grows.
 548
 549$$
 550o_i = \frac{\sum_{j \le i} \big(\phi(q_i) \cdot \phi(k_j)\big)\, v_j}{\sum_{j \le i} \phi(q_i) \cdot \phi(k_j)}
 551= \frac{\phi(q_i)^\top S_i}{\phi(q_i)^\top z_i},
 552\qquad S_i = S_{i-1} + \phi(k_i)\, v_i^\top,
 553\qquad z_i = z_{i-1} + \phi(k_i)
 554$$
 555
 556**Symbols**
 557
 558| Symbol | Meaning here | In the example |
 559|---|---|---|
 560| $o_i$ | the output for token $i$ | $o_3 = 4.13$ |
 561| $q_i$, $k_j$, $v_j$ | query of token $i$; key and value of token $j$ | $q_3 = (1, 0)$ |
 562| $\phi$ | the feature map, applied to each number: $x + 1$ if $x > 0$, else $e^x$; always positive so no weight is negative | φ(1, 0) = (2, 1) |
 563| $\sum_{j \le i}$ | add up over every token $j$ up to and including $i$ (causal) | $j$ = 1, 2, 3 |
 564| $\phi(q_i) \cdot \phi(k_j)$ | the unnormalized weight, a dot product in place of $e^{q \cdot k}$ | 5, 4, 6 |
 565| $v_i^\top$ | the value laid on its side as a row | |
 566| $\phi(k_i)\, v_i^\top$ | an **outer product**: a column times a row, giving a small table whose entry (m, c) is $\phi(k_i)_m \, v_{i,c}$ | (2, 1) × 2 = (4, 2) |
 567| $S_i$ | the pot after token $i$: the sum of those tables | (20, 22) |
 568| $z_i$ | the running sum of $\phi(k_j)$, for the denominator | (5, 5) |
 569| $\phi(q_i)^\top S_i$ | the query's ladle: its dot product with each column of $S_i$ | 62 |
 570
 571**In words:** "each output is a weighted average of the values so far, with
 572weights φ(q)·φ(k); because those weights split into a query part and a key
 573part, the key-and-value part can be kept as a running sum, and each query
 574reads that sum once."
 575
 576**With the numbers:** S₃ = (20, 22), z₃ = (5, 5), φ(q₃) = (2, 1):
 577o₃ = (40 + 22) / (10 + 5) = 62/15 = 4.133.
 578
 579**In Python:**
 580
 581```python
 582import math
 583def phi(v):
 584    # elu(x) + 1: x + 1 above zero, e^x at or below it
 585    return [x + 1 if x > 0 else math.exp(x) for x in v]
 586def dot(a, b):
 587    return sum(a_m * b_m for a_m, b_m in zip(a, b))
 588keys = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]
 589values = [2.0, 4.0, 6.0]
 590q3 = [1.0, 0.0]
 591[phi(k) for k in keys]  # → [[2.0, 1.0], [1.0, 2.0], [2.0, 2.0]]
 592# S: the running sum of φ(k_j) v_j; z: the running sum of φ(k_j)
 593S = [sum(phi(k)[m] * v for k, v in zip(keys, values)) for m in range(2)]
 594S  # → [20.0, 22.0]
 595z = [sum(phi(k)[m] for k in keys) for m in range(2)]
 596z  # → [5.0, 5.0]
 597# o_3 = φ(q)ᵀS / φ(q)ᵀz
 598round(dot(phi(q3), S) / dot(phi(q3), z), 3)  # → 4.133
 599# the quadratic form: the weights φ(q)·φ(k_j), then their weighted average
 600weights = [dot(phi(q3), phi(k)) for k in keys]
 601weights  # → [5.0, 4.0, 6.0]
 602round(dot(weights, values) / sum(weights), 3)  # → 4.133
 603# softmax attention on the same numbers (unscaled), for contrast
 604e = [math.exp(dot(q3, k)) for k in keys]
 605round(dot(e, values) / sum(e), 3)  # → 4.0
 606```
 607
 608![Side by side on the same queries and keys: softmax attention puts almost half of each late row on one key, while linear attention spreads each row thinly, its largest weight about 0.2](figures/primer.ml.efficient_architectures.linear_vs_softmax.svg)
 609
 610**Reading it:** the two heatmaps use the same ten queries and keys; rows are
 611queries, columns are keys, and darker means more weight. Both are causal
 612triangles and every row sums to 1. On the left, softmax picks favourites:
 613the exponential stretches gaps between scores (see `primer.ml.attention`), so
 614in the last five rows the largest weight averages 0.46. On the right, linear
 615attention spreads the same rows thinly: its largest weight averages 0.21. The
 616queries here are twice the usual size, and that is the telling part: at the
 617usual size the two numbers are 0.28 and 0.20, so making a query more decisive
 618sharpens softmax a lot and linear attention hardly at all, because φ(q)·φ(k)
 619never stretches a gap the way $e^x$ does. That bluntness is the price of the
 620pot: a fixed-size summary cannot pick out one exact token as sharply.
 621
 622**Why it matters in practice.** The work drops from n² × d to n × d², and
 623generation needs only the pot, not a cache. The catch is recall: asked to
 624copy back a specific token from far away, a fixed-size pot does worse than a
 625full cache. Later work added forgetting (decay) to the pot, which turns out
 626to be exactly the state-space models of the next section; the Mamba-2 paper
 627shows the two views describe the same computation.
 628
 629**In code:** `feature_map` is φ, `linear_attention_quadratic` builds all n × n weights (from `linear_attention_weights`) and blends the values, and `linear_attention_recurrent` gets the same outputs from the two running sums, returning the fixed-size S and z.
 630
 631## 5. State-space models: a summary with a fixed update rule
 632
 633### A recurrence that is also a convolution
 634
 635**Everyday picture.** A cup of tea cooling on a desk. Every minute it keeps
 636some fraction of its heat and gains whatever hot water you pour in. Its
 637temperature right now is a summary of everything you ever poured, with older
 638pours counting less. That is a recurrent network's one-page summary (see
 639`primer.ml.cnn_rnn`), with one difference: the rewrite rule is *linear*
 640(multiply and add, nothing more), and linearity buys a second way to compute
 641the same thing.
 642
 643**Tiny worked example.** The state keeps half of itself each step: A = 0.5,
 644B = 1, C = 1. Inputs x = (1, 0, 0, 2).
 645
 646| step t | x_t | h_t = 0.5 · h_(t−1) + x_t | y_t |
 647|---|---|---|---|
 648| 1 | 1 | 0.5 · 0 + 1 = 1 | 1 |
 649| 2 | 0 | 0.5 · 1 + 0 = 0.5 | 0.5 |
 650| 3 | 0 | 0.5 · 0.5 + 0 = 0.25 | 0.25 |
 651| 4 | 2 | 0.5 · 0.25 + 2 = 2.125 | 2.125 |
 652
 653Unroll the loop and each output is a weighted sum of all past inputs with
 654weights 1, 0.5, 0.25, 0.125 for "now, 1 step ago, 2 steps ago, 3 steps ago":
 655y₄ = 1·2 + 0.5·0 + 0.25·0 + 0.125·1 = 2.125. The same answer, with no loop.
 656
 657```mermaid
 658flowchart LR
 659  X["inputs x_1 ... x_T"] --> R["recurrence<br/>h_t = A h_(t−1) + B x_t<br/>one step at a time"]
 660  X --> K["convolution<br/>slide the kernel (CB, CAB, CA²B, ...)<br/>all positions at once"]
 661  R --> Y["the same outputs y_1 ... y_T"]
 662  K --> Y
 663```
 664
 665**Reading it:** two roads from the same inputs to the same outputs. The top
 666road is how the model *generates*: one token at a time, carrying a small
 667state h, with constant memory. The bottom road is how it *trains*: the
 668kernel is computed once, and every output is a weighted sum that can be done
 669in parallel, like a convolution in `primer.ml.cnn_rnn`. Classic RNNs only had
 670the top road, which is why they trained slowly.
 671
 672$$
 673h_t = A\,h_{t-1} + B\,x_t, \qquad y_t = C\,h_t
 674\qquad\Longleftrightarrow\qquad
 675y_t = \sum_{k=0}^{t-1} \big(C A^{k} B\big)\, x_{t-k}
 676$$
 677
 678**Symbols**
 679
 680| Symbol | Meaning here | In the example |
 681|---|---|---|
 682| $x_t$ | the input at step $t$ (one channel) | (1, 0, 0, 2) |
 683| $h_t$ | the state after step $t$: $N$ numbers (here $N$ = 1); $h_0 = 0$ | 1, 0.5, 0.25, 2.125 |
 684| $A$ | $N \times N$ matrix: how the old state carries over (its size sets how fast things fade) | 0.5 |
 685| $B$ | how the input is written into the state | 1 |
 686| $C$ | how the state is read out | 1 |
 687| $y_t$ | the output at step $t$ | |
 688| $\Longleftrightarrow$ | "the same thing, written another way" | |
 689| $A^k$ | $A$ multiplied by itself $k$ times; $A^0$ is "do nothing" | 0.5³ = 0.125 |
 690| $C A^k B$ | the **kernel**: how much an input $k$ steps ago still counts now | 1, 0.5, 0.25, 0.125 |
 691| $\sum_{k=0}^{t-1}$ | add over every look-back distance $k$ from 0 to $t - 1$ | |
 692
 693**In words:** "the state keeps a fraction of itself and adds the new input,
 694and the output reads the state; equivalently, the output is every past input
 695weighted by how much it has faded since."
 696
 697**With the numbers:** y₄ = C A⁰ B x₄ + C A¹ B x₃ + C A² B x₂ + C A³ B x₁ =
 6981·2 + 0.5·0 + 0.25·0 + 0.125·1 = 2.125.
 699
 700**In Python:**
 701
 702```python
 703A, B, C = 0.5, 1.0, 1.0
 704x = [1.0, 0.0, 0.0, 2.0]
 705h, y = 0.0, []
 706for x_t in x:
 707    # keep half of the old state, add the new input
 708    h = A * h + B * x_t
 709    y.append(C * h)
 710y  # → [1.0, 0.5, 0.25, 2.125]
 711# the kernel C A^k B, for k = 0, 1, 2, 3
 712kernel = [C * A ** k * B for k in range(4)]
 713kernel  # → [1.0, 0.5, 0.25, 0.125]
 714# y_4 as a weighted sum of all four inputs, newest first
 715sum(kernel[k] * x[3 - k] for k in range(4))  # → 2.125
 716```
 717
 718![Three kernels: with A = 0.5 an input is forgotten within about 5 steps, with 0.9 within about 40, with 0.99 it still counts about a fifth after 150 steps](figures/primer.ml.efficient_architectures.ssm_kernel.svg)
 719
 720**Reading it:** the x-axis is how many steps ago an input arrived; the
 721y-axis is how much it still counts (the kernel C A^k B). Each curve is one
 722value of A. With A = 0.5 the curve collapses almost at once: short memory.
 723With A = 0.99 it is still well above zero after 150 steps: long memory. A
 724real SSM layer has thousands of channels, each with its own A, so it
 725remembers at many time scales at once. S4 was the first to make this work on
 726very long sequences, by choosing and computing these kernels carefully.
 727
 728**In code:** `ssm_recurrent` runs the loop, `ssm_kernel` builds C A^k B, and `ssm_convolution` gets the same outputs from the kernel with a single NumPy convolution.
 729
 730### Selective: letting each token decide what to keep (Mamba)
 731
 732**Everyday picture.** A note-taker with a dial. For filler words ("um", "so",
 733"anyway") they barely touch their notes. For a name or a number they wipe the
 734relevant line and write the new fact. A fixed SSM uses the same dial setting
 735for every word, so it must either write everything (and forget quickly) or
 736write little (and never take in the important word properly). Mamba reads the
 737dial setting off each token itself: that is what *selective* means.
 738
 739The dial is a step size Δ. The parameters of the state update are derived
 740from it at every step (this is called *discretization*: turning a continuous
 741rate of change into one step's keep and write amounts):
 742
 743- keep factor $\bar{A} = e^{\Delta a}$, with $a$ negative, so a big Δ makes
 744  it nearly 0 (forget) and a tiny Δ nearly 1 (keep);
 745- write factor $\bar{B}$, which goes the opposite way.
 746
 747**Tiny worked example.** One channel, a = −1, b = 1. A token that carries a
 748marker gets Δ ≈ 5; an ordinary token gets Δ ≈ 0.0067.
 749
 750| token | Δ | keep Ā = e^(−Δ) | write B̄ = 1 − e^(−Δ) | effect |
 751|---|---|---|---|---|
 752| marked | 5.007 | 0.0067 | 0.9933 | replace the state with this token |
 753| ordinary | 0.0067 | 0.9933 | 0.0067 | leave the state almost untouched |
 754
 755The recall task: a sequence of small noise values with one marked 7 in third
 756place, then nine more noise values. The selective SSM writes the 7 almost
 757fully and then keeps it: at the end its state is **6.55**. With one fixed Δ
 758for every token, the best any setting manages is **0.26**: a Δ big enough to
 759write the 7 also lets the nine later tokens overwrite it.
 760
 761```mermaid
 762flowchart LR
 763  U["token u_t"] --> D["Δ_t = softplus(w · u_t + β)<br/>how much this token matters"]
 764  D --> AB["Ā_t = e^(Δ_t a): keep<br/>B̄_t: write"]
 765  U --> XV["x_t: what to write"]
 766  AB --> H["h_t = Ā_t h_(t−1) + B̄_t x_t"]
 767  XV --> H
 768  HP["h_(t−1)"] --> H
 769  H --> Y["y_t = C h_t"]
 770```
 771
 772**Reading it:** the new part is the top path. The token itself feeds a small
 773function that produces its step size Δ, which sets how much of the old state
 774to keep and how much of this token to write. In Mamba the matrices B and C
 775are also computed from the token, by the same kind of path. Everything below
 776is the recurrence from before, except that Ā and B̄ now change at every step.
 777
 778$$
 779\Delta_t = \operatorname{softplus}(w\,u_t + \beta), \qquad
 780\bar{A}_t = e^{\Delta_t a}, \qquad
 781\bar{B}_t = \frac{e^{\Delta_t a} - 1}{a}\, b, \qquad
 782h_t = \bar{A}_t\, h_{t-1} + \bar{B}_t\, x_t
 783$$
 784
 785**Symbols**
 786
 787| Symbol | Meaning here | In the example |
 788|---|---|---|
 789| $u_t$ | what the model can see about token $t$ (here, just its marker, 0 or 1) | 1 for the 7 |
 790| $w$, $\beta$ | learned weight and offset that turn $u_t$ into a step size | 10, −5 |
 791| softplus | $\log(1 + e^x)$: a smooth ramp that is always positive, about $x$ for large $x$ and about 0 for very negative $x$ | softplus(5) = 5.007 |
 792| $\Delta_t$ | the step size for token $t$: how much this token matters | 5.007 or 0.0067 |
 793| $a$ | a negative learned rate; more negative means faster forgetting | −1 |
 794| $b$ | how strongly inputs are written | 1 |
 795| $\bar{A}_t$ | this step's keep factor ("A-bar") | 0.0067 or 0.9933 |
 796| $\bar{B}_t$ | this step's write factor ("B-bar") | 0.9933 or 0.0067 |
 797| $x_t$, $h_t$ | the value written, and the state after step $t$ | 7, then 6.95 |
 798
 799**In words:** "each token computes how much it matters; that sets how much
 800of the old state survives and how much of the token gets written; then the
 801usual update runs with those per-token amounts."
 802
 803**With the numbers:** the marked 7 has Δ = softplus(10·1 − 5) = 5.007, so
 804Ā = e^(−5.007) = 0.0067 and B̄ = (0.0067 − 1)/(−1) · 1 = 0.9933: the state
 805becomes about 0.9933 × 7 = 6.95. Each ordinary token after it keeps 0.9933 of
 806the state, and nine of them leave about 6.95 × 0.9933⁹ = 6.55.
 807
 808**In Python:**
 809
 810```python
 811import math
 812def softplus(x):
 813    return math.log(1 + math.exp(x))
 814a, b = -1.0, 1.0
 815# Δ for a marked token, then for an ordinary one
 816d_mark, d_plain = softplus(10 * 1 - 5), softplus(10 * 0 - 5)
 817round(d_mark, 3), round(d_plain, 4)  # → (5.007, 0.0067)
 818# Ā_t and B̄_t for the marked token: keep almost nothing, write almost everything
 819keep_mark, write_mark = math.exp(d_mark * a), (math.exp(d_mark * a) - 1) / a * b
 820round(keep_mark, 4), round(write_mark, 4)  # → (0.0067, 0.9933)
 821# and for an ordinary token: keep almost everything, write almost nothing
 822keep_plain = math.exp(d_plain * a)
 823round(keep_plain, 4)  # → 0.9933
 824# the 7 is written once, then kept through nine ordinary tokens
 825round(write_mark * 7 * keep_plain ** 9, 2)  # → 6.55
 826```
 827
 828![The selective state jumps to about 7 at the marked token and holds near 6.5 to the end, while the best fixed step peaks near 0.6 and ends at 0.26, and a large fixed step just copies the latest noise](figures/primer.ml.efficient_architectures.selective_recall.svg)
 829
 830**Reading it:** the x-axis is the position in the sequence; grey bars are
 831the input values, with the marked 7 at position 2. The blue line is the
 832selective SSM's state: it jumps to the 7 and then barely moves, because
 833every later token has a tiny Δ. The red line is the best fixed Δ: it can
 834only write a small fraction of each input, so the 7 lifts it to about 0.6 at
 835most and it ends at 0.26. The
 836orange line is a large fixed Δ: it writes every token fully, so it tracks
 837whatever arrived last and the 7 is gone by the next step. Only the selective
 838model can both write the important token and ignore the rest.
 839
 840**In code:** `discretize` turns Δ into Ā and B̄, `softplus` and `step_sizes` compute each token's Δ, `selective_ssm` runs the per-token recurrence, and `selective_recall` and `time_invariant_recall` run the recall task.
 841
 842### Training in parallel when the kernel keeps changing: the scan
 843
 844**Everyday picture.** A relay race where each runner must know the total
 845time so far. Done one after another it takes as long as the whole race. But
 846"multiply by a, then add b" steps can be *merged*: two consecutive steps are
 847themselves one step of the same shape. So pairs of runners merge their legs,
 848then pairs of pairs, like a knockout tournament: log₂ n rounds instead of n (log₂ n is how many
 849times n can be halved before reaching 1: 10 for 1,024).
 850
 851**Tiny worked example.** The four steps of the tea example are
 852(a, b) = (0.5, 1), (0.5, 0), (0.5, 0), (0.5, 2), meaning "h becomes a·h + b".
 853Round 1: every step merges with the one before it. Round 2: every step merges
 854with the result two places before it. After 2 rounds (log₂ 4 = 2), the b
 855parts are 1, 0.5, 0.25, 2.125: every state from the loop.
 856
 857```mermaid
 858flowchart TB
 859  s1["(0.5, 1)"]
 860  s2["(0.5, 0)"]
 861  s3["(0.5, 0)"]
 862  s4["(0.5, 2)"]
 863  s1 --> r2["(0.25, 0.5)"]
 864  s2 --> r2
 865  s2 --> r3["(0.25, 0)"]
 866  s3 --> r3
 867  s3 --> r4["(0.25, 2)"]
 868  s4 --> r4
 869  s1 --> f3["(0.125, 0.25)"]
 870  r3 --> f3
 871  r2 --> f4["(0.0625, 2.125)"]
 872  r4 --> f4
 873```
 874
 875**Reading it:** the top row holds the four steps. Each arrow pair is one
 876merge; every merge in a row happens at the same time. The middle row is
 877round 1, the bottom row round 2. Read the second number in each final box
 878(and in s1 and r2, which were already complete): 1, 0.5, 0.25, 2.125, the
 879states from the table above. A selective SSM cannot use the convolution road,
 880because its Ā changes every step, but it can use this one.
 881
 882$$
 883(a_1, b_1) \circ (a_2, b_2) = \big(a_1 a_2,\; a_2 b_1 + b_2\big)
 884$$
 885
 886**Symbols**
 887
 888| Symbol | Meaning here | In the example |
 889|---|---|---|
 890| $(a, b)$ | one step: "multiply the state by $a$, then add $b$" | (0.5, 1) |
 891| $\circ$ | "do the first step, then the second": merging two steps into one | |
 892| $a_1 a_2$ | the combined multiplier | 0.25 |
 893| $a_2 b_1 + b_2$ | the first step's addition, shrunk by the second step, plus the second's own | 0.5·1 + 0 = 0.5 |
 894
 895**In words:** "doing step 1 then step 2 is the same as one step that
 896multiplies by both and adds step 1's contribution, faded by step 2."
 897
 898**With the numbers:** (0.25, 0.5) ∘ (0.25, 2) = (0.0625, 0.25·0.5 + 2) =
 899(0.0625, 2.125), the last box in the diagram.
 900
 901**In Python:**
 902
 903```python
 904def merge(first, then):
 905    a1, b1 = first
 906    a2, b2 = then
 907    return (a1 * a2, a2 * b1 + b2)
 908steps = [(0.5, 1.0), (0.5, 0.0), (0.5, 0.0), (0.5, 2.0)]
 909# round 1: each step absorbs the one just before it
 910r1 = [steps[0]] + [merge(steps[t - 1], steps[t]) for t in range(1, 4)]
 911r1  # → [(0.5, 1.0), (0.25, 0.5), (0.25, 0.0), (0.25, 2.0)]
 912# round 2: each absorbs the result two places before it
 913r2 = r1[:2] + [merge(r1[t - 2], r1[t]) for t in range(2, 4)]
 914[b for a, b in r2]  # → [1.0, 0.5, 0.25, 2.125]
 915```
 916
 917**Why it matters in practice.** This is how Mamba trains on long sequences
 918at GPU speed despite being recurrent: a parallel scan, written so the state
 919stays in fast on-chip memory. At generation time it switches back to the
 920plain loop, one token at a time, with a state that never grows.
 921
 922**In code:** `parallel_scan` runs the rounds for any length and returns the states and the number of rounds: 10 for 1,024 steps.
 923
 924### Constant memory, and hybrids
 925
 926**Everyday picture.** An SSM travels with a backpack of fixed size; the KV
 927cache is a suitcase that grows with every token. A *hybrid* model mostly uses
 928backpacks but brings one suitcase for every few layers, so it can still look
 929up an exact earlier token when it needs to.
 930
 931**Tiny worked example.** One SSM layer with 4,096 channels and a 16-number
 932state per channel holds 4,096 × 16 × 2 bytes = **128 KB**, at 10 tokens or at
 93310 million. One attention layer of the running 8B example holds **537 MB** at
 934128k tokens. Build 32 layers as 4 attention layers and 28 SSM layers (one in
 935eight, the ratio Jamba uses) and the cache at 128k drops from 17.2 GB to
 936**2.15 GB**.
 937
 938```mermaid
 939flowchart TB
 940  I["tokens in"] --> M1["SSM layer<br/>fixed state"] --> M2["SSM layer"] --> M3["..."] --> A1["attention layer<br/>KV cache grows"]
 941  A1 --> M4["SSM layer"] --> M5["..."] --> A2["attention layer"] --> O["next-token scores"]
 942```
 943
 944**Reading it:** most boxes are SSM layers, each carrying a fixed-size state.
 945Every so often an attention layer sits in the stack; only those keep a KV
 946cache. The memory bill therefore scales with the number of attention layers,
 947not with the total depth, while the attention layers are still there for
 948exact copying and lookup.
 949
 950$$
 951\text{bytes} = L_{\text{att}} \cdot n \cdot 2\,H_{kv}\,d_h\,b \;+\; (L - L_{\text{att}}) \cdot D \cdot N \cdot b
 952$$
 953
 954**Symbols**
 955
 956| Symbol | Meaning here | In the example |
 957|---|---|---|
 958| $L$ | total layers | 32 |
 959| $L_{\text{att}}$ | how many of them are attention layers | 4 |
 960| $n$ | context length | 131,072 |
 961| $2\,H_{kv}\,d_h\,b$ | one attention layer's cache per token (from section 1) | 4,096 bytes |
 962| $D$ | channels in an SSM layer | 4,096 |
 963| $N$ | state numbers per channel | 16 |
 964| $b$ | bytes per number | 2 |
 965
 966**In words:** "the attention layers pay per token as before; the SSM layers
 967pay a fixed amount that does not depend on n at all."
 968
 969**With the numbers:** 4 × 131,072 × 4,096 = 2,147,483,648 bytes, plus
 97028 × 4,096 × 16 × 2 = 3,670,016 bytes: 2.15 GB, about one eighth of 17.2 GB.
 971
 972**In Python:**
 973
 974```python
 975n, H_kv, d_h, b = 131_072, 8, 128, 2
 976D, N = 4096, 16
 977# one SSM layer's state, in KB, at any context length
 978D * N * b / 1024  # → 128.0
 979# one attention layer's KV cache at 128k tokens, in MB
 980n * 2 * H_kv * d_h * b / 1e6  # → 536.870912
 981# 32 layers: all attention, then 4 attention + 28 SSM, in GB
 982round(32 * n * 2 * H_kv * d_h * b / 1e9, 2)  # → 17.18
 983round((4 * n * 2 * H_kv * d_h * b + 28 * D * N * b) / 1e9, 2)  # → 2.15
 984```
 985
 986**Why it matters in practice.** Pure SSMs are strong at language modelling
 987but weaker than attention at copying and exact recall over long contexts,
 988for the same reason as linear attention: a fixed-size state is a summary.
 989Hybrids such as Jamba keep a few attention layers for those jobs and get most
 990of the SSM's memory savings.
 991
 992**In code:** `ssm_state_bytes` is the fixed backpack, and `hybrid_cache_bytes` adds up a mixed stack.
 993
 994## 6. Compressing the KV cache
 995
 996### Fewer key/value heads: GQA and MQA, a recap
 997
 998**Everyday picture.** Colleagues sharing one reference binder instead of each
 999keeping a personal copy: everyone still asks their own questions, but the
1000shelf holds fewer binders.
1001
1002**Tiny worked example.** On the running example (32 layers, 128 numbers per
1003head, 16-bit), the cache per token for different numbers of KV heads:
1004
1005| Design | KV heads | Cache per token |
1006|---|---|---|
1007| multi-head attention | 32 | 512 KB |
1008| grouped-query attention | 8 | 128 KB |
1009| multi-query attention | 1 | 16 KB |
1010
1011The diagram of shared heads and the full memory arithmetic live in
1012`primer.ml.attention` and `primer.ml.inference`; everything below stacks on
1013top of whichever of these a model uses.
1014
1015### Latent KV: cache the ingredients, cook on demand
1016
1017**Everyday picture.** A restaurant that stores ingredients, not finished
1018dishes. Every head's keys and values can be cooked from a short list of
1019ingredients (the *latent* vector) with a fixed recipe shared by every token.
1020Store the ingredients, keep the recipe once, and cook when a query needs it.
1021Better still, the recipe can be folded into the query itself, so nothing is
1022ever cooked at all.
1023
1024**Tiny worked example.** A token x = (1, 0, 2, 1), two heads of width 2.
1025Ordinary attention would cache 2 heads × 2 numbers × (key and value) = 8
1026numbers. Instead:
1027
10281. Squeeze: c = x W_down = (1 + 2, 0 + 1) = **(3, 1)**. Only these 2 numbers
1029   are cached: 4 times smaller.
10302. When needed, expand: k = c W_uk = **(3, 1, 3, 1)**, both heads' keys.
10313. Or fold the expansion into the query: for q = (1, 2, 0, 1), q · k = 6, and
1032   (q W_ukᵀ) · c = (1, 3) · (3, 1) = 6. Same score, and k was never built.
1033
1034DeepSeek-V2 uses this at scale: its 128 heads of width 128 would cache 32,768
1035numbers per token per layer; its latent caches 512, plus 64 for a small key
1036that carries position information: **576, about 57 times less**.
1037
1038```mermaid
1039flowchart LR
1040  X["token x<br/>(d_model numbers)"] -->|"W_down"| C[("cache: latent c<br/>d_c numbers")]
1041  C -->|"W_uk"| K["keys for every head"]
1042  C -->|"W_uv"| V["values for every head"]
1043  Qn["query q"] -->|"absorbed: q W_ukᵀ"| QL["query in latent space"]
1044  QL -->|"score directly against c"| C
1045```
1046
1047**Reading it:** the cylinder is all that is stored per token. The two arrows
1048to the right show the plain way: rebuild keys and values from the latent.
1049The bottom path shows the absorbed way: translate the query into latent space
1050once, then score it against the cached latents directly, and mix latents
1051before expanding with W_uv. Both paths give identical outputs; the absorbed
1052one never materializes the big keys and values.
1053
1054$$
1055c_t = x_t W^{D}, \qquad k_t = c_t W^{UK}, \qquad v_t = c_t W^{UV},
1056\qquad q \cdot k_t = \big(q\, {W^{UK}}^{\top}\big) \cdot c_t
1057$$
1058
1059**Symbols**
1060
1061| Symbol | Meaning here | Shape / example |
1062|---|---|---|
1063| $x_t$ | token $t$'s vector coming into the layer | $d_\text{model}$; (1, 0, 2, 1) |
1064| $W^{D}$ | the learned "down" projection that squeezes | $d_\text{model} \times d_c$ |
1065| $c_t$ | the latent: the only thing cached | $d_c$; (3, 1) |
1066| $W^{UK}$, $W^{UV}$ | learned "up" projections to every head's keys and values | $d_c \times H d_h$ |
1067| $k_t$, $v_t$ | token $t$'s keys and values for all heads, side by side | $H d_h$; $k$ = (3, 1, 3, 1) |
1068| $q$ | a query (one head's slice, or all heads side by side) | (1, 2, 0, 1) |
1069| ${W^{UK}}^{\top}$ | $W^{UK}$ transposed (rows become columns) | $H d_h \times d_c$ |
1070| $q\,{W^{UK}}^{\top}$ | the query translated into latent space | $d_c$; (1, 3) |
1071
1072**In words:** "squeeze each token into a short latent and cache only that;
1073keys and values are the latent times fixed up-projections, so a query can
1074be moved into latent space once and scored against the cache directly."
1075
1076**With the numbers:** c = (3, 1), k = (3, 1, 3, 1), q · k = 3 + 2 + 0 + 1 =
10776, and (1, 3) · (3, 1) = 3 + 3 = 6.
1078
1079**In Python:**
1080
1081```python
1082x = [1, 0, 2, 1]
1083W_down = [[1, 0], [0, 1], [1, 0], [0, 1]]
1084W_uk = [[1, 0, 1, 0], [0, 1, 0, 1]]
1085def vecmat(v, M):
1086    # row vector times matrix: entry c is Σ_r v_r M[r][c]
1087    return [sum(v[r] * M[r][c] for r in range(len(v))) for c in range(len(M[0]))]
1088# c = x W_down: the only thing cached
1089c = vecmat(x, W_down)
1090c  # → [3, 1]
1091# k = c W_uk: both heads' keys, rebuilt on demand
1092k = vecmat(c, W_uk)
1093k  # → [3, 1, 3, 1]
1094q = [1, 2, 0, 1]
1095sum(q_m * k_m for q_m, k_m in zip(q, k))  # → 6
1096# absorbed: move q into latent space, then score against c
1097W_uk_T = [list(col) for col in zip(*W_uk)]
1098q_latent = vecmat(q, W_uk_T)
1099q_latent, sum(a * b for a, b in zip(q_latent, c))  # → ([1, 3], 6)
1100# DeepSeek-V2's shape: cached numbers per token per layer
11012 * 128 * 128, 512 + 64, round(2 * 128 * 128 / (512 + 64), 1)  # → (32768, 576, 56.9)
1102```
1103
1104Why can so few numbers stand in for so many? Because every head's keys and
1105values are built from the same token, they are highly redundant. In the
1106language of `primer.notation`, the key projection W^D W^UK has **low rank**:
1107however many numbers it outputs, they all vary along only d_c independent
1108directions. Storing those d_c coordinates loses nothing that projection can
1109produce. One wrinkle: rotary position embeddings (see `primer.ml.positional`)
1110rotate each key by its position, which breaks the absorption trick, so
1111DeepSeek-V2 carries position in that separate small 64-number key.
1112
1113![Cache per token on the 32-layer example: 512 KB for 32 KV heads, 128 KB for 8, 36 KB for a 576-number latent, 33 KB for 8 heads at 4 bits, 16 KB for one head](figures/primer.ml.efficient_architectures.kv_per_token.svg)
1114
1115**Reading it:** each bar is the cache one token costs across all 32 layers
1116of the running example, for one design. The top bar is classic multi-head
1117attention; every bar below it is a way of shrinking it. Sharing heads (GQA,
1118MQA), squeezing into a latent, and cutting bits land in the same range, tens
1119of kilobytes instead of hundreds, by different routes that can be combined.
1120
1121**In code:** `LatentKVAttention` holds the four projections; `LatentKVAttention.compress` produces the cache, `LatentKVAttention.attend` expands it into keys and values, and `LatentKVAttention.attend_absorbed` gets the identical output without expanding. `latent_kv_worked_example` is the (3, 1) example and `latent_kv_bytes_per_token` the memory arithmetic.
1122
1123### Quantizing the cache
1124
1125**Everyday picture.** The same move as for weights in `primer.ml.inference`:
1126write each number to the nearest tenth instead of the nearest thousandth,
1127with one ruler per stored vector so one big number does not coarsen all the
1128others.
1129
1130**Tiny worked example.** A cached key (0.7, −0.3, 0.2, 0.04) at 4 bits (codes
1131−7 to 7). The step size is 0.7 / 7 = 0.1, the codes are **(7, −3, 2, 0)**, and
1132reading back gives (0.7, −0.3, 0.2, 0): the tiny 0.04 is lost. A 128-number
1133head vector drops from 256 bytes to 64 bytes of codes plus a 2-byte scale:
113466 bytes, **3.9 times smaller**.
1135
1136```mermaid
1137flowchart LR
1138  KV["new key or value<br/>16-bit numbers"] -->|"scale = max / 7<br/>code = round(x / scale)"| ST[("cache: 4-bit codes<br/>+ one scale per vector")]
1139  ST -->|"code × scale"| R["approximate key or value"]
1140  R --> A["attention as usual"]
1141```
1142
1143**Reading it:** writing to the cache quantizes once per token; reading from
1144it multiplies back. Attention itself is unchanged; it just sees slightly
1145rounded keys and values. The saving is on the cylinder, the part that grows
1146with every token.
1147
1148$$
1149\text{bytes per token} = 2\,L\,H_{kv}\left(d_h \cdot \frac{\text{bits}}{8} + \frac{\text{scale bits}}{8}\right)
1150$$
1151
1152**Symbols**
1153
1154| Symbol | Meaning here | In the example |
1155|---|---|---|
1156| $2\,L\,H_{kv}$ | how many head vectors one token stores: a key and a value per KV head per layer | 2 × 32 × 8 = 512 |
1157| $d_h$ | numbers per head vector | 128 |
1158| bits / 8 | bytes per code | 4-bit: 0.5 |
1159| scale bits / 8 | bytes for the one scale each vector carries | 16-bit: 2 |
1160
1161**In words:** "every stored head vector costs its codes plus one scale, and
1162a token stores a key vector and a value vector per head per layer."
1163
1164**With the numbers:** 2 × 32 × 8 × (128 × 0.5 + 2) = 512 × 66 = 33,792
1165bytes, against 131,072 at 16 bits: 3.9 times smaller.
1166
1167**In Python:**
1168
1169```python
1170k = [0.7, -0.3, 0.2, 0.04]
1171# one scale per vector: the largest magnitude lands on code 7
1172s = max(abs(k_j) for k_j in k) / 7
1173round(s, 3)  # → 0.1
1174codes = [round(k_j / s) for k_j in k]
1175codes  # → [7, -3, 2, 0]
1176[round(s * c, 2) for c in codes]  # → [0.7, -0.3, 0.2, 0.0]
1177L, H_kv, d_h = 32, 8, 128
1178# bytes per token: 16-bit, then 4-bit codes plus a 16-bit scale per vector
11792 * L * H_kv * d_h * 16 / 8, 2 * L * H_kv * (d_h * 4 / 8 + 16 / 8)  # → (131072.0, 33792.0)
1180```
1181
1182**Why it matters in practice.** On random test data, an 8-bit cache moves
1183the attention output by under 1% and a 4-bit cache by about 12%; trained
1184models tolerate this far better than random data suggests, and careful
1185schemes go lower. Keys tend to have a few channels with consistently large
1186values, so KIVI quantizes keys per channel and values per token and reaches
11872 bits. Quantization stacks with everything above: a GQA model with a sliding
1188window and a 4-bit cache enjoys all three savings.
1189
1190**In code:** `quantize_kv` quantizes each cached vector with its own scale via `primer.ml.inference.quantize`, `quantized_kv_bytes_per_token` is the formula, and `quantized_cache_error` measures how far attention's output moves.
1191
1192## Putting it together
1193
1194![On log axes from 1k to 1M tokens, multi-head and grouped-query caches climb past 80 GB, latent and 4-bit caches climb 4 times lower, the hybrid 8 times lower, while the sliding window flattens at 0.5 GB and the SSM state stays at 4 MB](figures/primer.ml.efficient_architectures.memory_vs_context.svg)
1195
1196**Reading it:** the x-axis is context length and the y-axis is cache memory
1197for one conversation, both logarithmic, with the dashed line at 80 GB. Lines
1198with the same slope grow the same way (in proportion to n); compression moves
1199a line down without changing its slope. The sliding window bends flat at
12004,096 tokens, and the SSM is flat from the start. Flat lines are what make
1201million-token contexts affordable; the price is that they no longer keep
1202every token exactly.
1203
1204| Technique | What it cuts | What it gives up |
1205|---|---|---|
1206| Sliding window | compute to n·w, cache to w tokens | direct access beyond the window |
1207| Sparse (global, strided) | compute | exact long-range pairs not in the pattern |
1208| Linear attention | compute to n·d², cache to a fixed state | sharp, exact recall |
1209| SSM / Mamba | the same | the same, softened by selectivity |
1210| Hybrid | most of the cache | a little of both |
1211| GQA / MQA | cache by the sharing factor | a little modelling capacity |
1212| Latent KV | cache by d_c / (2·H·d_h) | extra matrix work, care with positions |
1213| Quantized KV | cache by 16 / bits | a little precision |
1214
1215**In code:** `cache_bytes_by_method` computes every line in the figure for a given context length.
1216
1217## In 20 seconds
1218
1219- Long context has two bills: attention scores grow with n², and the KV
1220  cache grows with n (per layer, per conversation). At 1M tokens the cache
1221  of an 8B-class model alone is about 137 GB.
1222- **Sliding window:** each token reads the last w; stacked layers still reach
1223  L·(w − 1) back, and a rolling buffer caps the cache at w tokens.
1224  **Sparse patterns** add a few global or strided links so any two tokens are
1225  a hop or two apart.
1226- **Linear attention** replaces e^(q·k) with φ(q)·φ(k), which turns attention
1227  into a running sum: O(n) work and a fixed-size state, but blurrier recall.
1228- **State-space models** update a fixed state linearly; fixed ones train as a
1229  convolution, **selective** ones (Mamba) let each token set how much to keep
1230  and write, and train with a parallel scan. **Hybrids** keep a few attention
1231  layers for exact lookup.
1232- **Compress the cache:** share KV heads (GQA/MQA), cache a small latent and
1233  expand on demand (latent KV), and store fewer bits per number.
1234
1235## Self-test questions
1236
1237**Why does doubling the context quadruple attention's compute but only
1238double its cache?**
1239Every token's query is scored against every key, so the scores form an
1240n × n table: doubling n quadruples it. The cache stores one key and one value
1241per token per layer, a list that grows by one entry per token, so doubling n
1242doubles it.
1243
1244**A model uses a 4,096-token sliding window in all 32 layers. Can token
1245100,000 be influenced by token 1?**
1246Yes, in principle: information moves up to w − 1 = 4,095 positions per
1247layer, so 32 layers reach 131,040 positions back. In practice it must be
1248relayed through about 25 intermediate tokens and layers, so it arrives
1249weakened; direct, exact lookup only works within the window.
1250
1251**What does a global token do in a sparse pattern, and why is it cheap?**
1252Every token reads it and it reads every token, so any two tokens are at most
1253two hops apart. It adds only about one column and one row of scores, a cost
1254that grows with n rather than n².
1255
1256**Why can linear attention run as a recurrence but softmax attention cannot?**
1257Linear attention's weight φ(q)·φ(k) splits into a query part and a key part,
1258so the key-and-value parts can be summed ahead of time into a fixed-size
1259state that any later query can read. The softmax weight e^(q·k) does not
1260split that way, so each new query must revisit every stored key.
1261
1262**What does "selective" mean in Mamba, and what problem does it fix?**
1263The step size Δ, and with it how much of the old state is kept and how much
1264of the new token is written, is computed from each token. A fixed SSM applies
1265the same keep and write amounts to every token, so it cannot both absorb one
1266important token and ignore the filler around it.
1267
1268**If a selective SSM's parameters change every step, how does it train in
1269parallel?**
1270The update "multiply by a, add b" can be merged: two consecutive steps form
1271one step of the same kind. A parallel scan merges pairs, then pairs of
1272pairs, and finishes in about log₂ n rounds.
1273
1274**How does latent KV caching save memory without changing the attention
1275outputs?**
1276Keys and values are computed as a small cached latent times fixed
1277up-projection matrices, so storing the latent is enough to rebuild them
1278exactly. The up-projection can even be folded into the query and output
1279side, so the full keys and values are never built.
1280
1281**Why do hybrid models keep a few attention layers instead of going all-SSM?**
1282A fixed-size state is a lossy summary, which hurts copying and exact recall
1283over long contexts. A handful of attention layers restores exact lookup
1284while most layers keep constant memory, so the cache shrinks roughly by the
1285fraction of layers that are SSMs.
1286
1287## The papers behind this lesson
1288
1289- **Child, Gray, Radford & Sutskever, *Generating Long Sequences with Sparse
1290  Transformers* (2019)**: https://arxiv.org/abs/1904.10509. Introduced
1291  strided and fixed sparse attention patterns, cutting attention's cost to
1292  about n√n.
1293- **Beltagy, Peters & Cohan, *Longformer: The Long-Document Transformer*
1294  (2020)**: https://arxiv.org/abs/2004.05150. Combined a sliding window with
1295  task-chosen global tokens to read long documents in linear time.
1296- **Zaheer et al., *Big Bird: Transformers for Longer Sequences* (2020)**:
1297  https://arxiv.org/abs/2007.14062. Mixed window, global and random links, and
1298  proved such sparse attention keeps the expressive power of full attention.
1299- **Katharopoulos et al., *Transformers are RNNs: Fast Autoregressive
1300  Transformers with Linear Attention* (2020)**:
1301  https://arxiv.org/abs/2006.16236. Replaced softmax with a kernel feature
1302  map, turning attention into a running sum with O(n) cost.
1303- **Gu, Goel & Ré, *Efficiently Modeling Long Sequences with Structured State
1304  Spaces* (S4, 2021)**: https://arxiv.org/abs/2111.00396. Made linear
1305  state-space layers trainable on very long sequences through the
1306  convolution view.
1307- **Gu & Dao, *Mamba: Linear-Time Sequence Modeling with Selective State
1308  Spaces* (2023)**: https://arxiv.org/abs/2312.00752. Made the state-space
1309  parameters depend on the input and trained them with a hardware-aware
1310  parallel scan.
1311  [Annotated companion](../../papers/mamba.html)
1312- **Dao & Gu, *Transformers are SSMs* (Mamba-2, 2024)**:
1313  https://arxiv.org/abs/2405.21060. Showed that selective SSMs and a form of
1314  linear attention are two views of one computation.
1315- **Jiang et al., *Mistral 7B* (2023)**: https://arxiv.org/abs/2310.06825.
1316  Used sliding-window attention with a rolling-buffer cache in a strong open
1317  model.
1318  [Annotated companion](../../papers/mistral-7b.html)
1319- **Xiao et al., *Efficient Streaming Language Models with Attention Sinks*
1320  (2023)**: https://arxiv.org/abs/2309.17453. Found that models lean on the
1321  first few tokens, and that keeping them plus a window allows streaming
1322  without limit.
1323- **Lieber et al., *Jamba: A Hybrid Transformer-Mamba Language Model*
1324  (2024)**: https://arxiv.org/abs/2403.19887. Interleaved one attention layer
1325  per seven Mamba layers to cut the KV cache while keeping recall.
1326- **Shazeer, *Fast Transformer Decoding: One Write-Head is All You Need*
1327  (2019)**: https://arxiv.org/abs/1911.02150. Introduced multi-query
1328  attention, one shared key/value head for all query heads.
1329  [Annotated companion](../../papers/multi-query-attention.html)
1330- **DeepSeek-AI, *DeepSeek-V2* (2024)**: https://arxiv.org/abs/2405.04434.
1331  Introduced multi-head latent attention, caching a small latent per token
1332  instead of full keys and values.
1333- **Liu et al., *KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV
1334  Cache* (2024)**: https://arxiv.org/abs/2402.02750. Quantized keys per
1335  channel and values per token, bringing the cache down to 2 bits.
1336
1337## Further reading
1338
1339- Child et al., *Sparse Transformers* (2019): https://arxiv.org/abs/1904.10509
1340- Beltagy et al., *Longformer* (2020): https://arxiv.org/abs/2004.05150
1341- Katharopoulos et al., *Transformers are RNNs* (2020): https://arxiv.org/abs/2006.16236
1342- Gu, Goel & Ré, *S4* (2021): https://arxiv.org/abs/2111.00396
1343- Sasha Rush et al., *The Annotated S4* (the S4 paper, line by line in code): https://srush.github.io/annotated-s4/
1344- Gu & Dao, *Mamba* (2023): https://arxiv.org/abs/2312.00752
1345- Dao & Gu, *Mamba-2* (2024): https://arxiv.org/abs/2405.21060
1346- Lieber et al., *Jamba* (2024): https://arxiv.org/abs/2403.19887
1347- DeepSeek-AI, *DeepSeek-V2* (2024): https://arxiv.org/abs/2405.04434
1348- Ainslie et al., *GQA* (2023): https://arxiv.org/abs/2305.13245
1349"""
1350
1351from __future__ import annotations
1352
1353from collections import deque
1354
1355import numpy as np
1356
1357from primer._show import banner, matrix, say, table, takeaway
1358from primer.ml.attention import causal_mask, scaled_dot_product_attention, softmax
1359from primer.ml.inference import dequantize, kv_cache_bytes, kv_cache_bytes_per_token, quantize
1360
1361# The running example: a Llama-3-8B-shaped model, the same one `primer.ml.inference` sizes.
1362LAYERS, KV_HEADS, HEAD_DIM = 32, 8, 128
1363
1364# ---------------------------------------------------------------------------
1365# 1. The bill: pairs scored and bytes cached
1366# ---------------------------------------------------------------------------
1367
1368
1369def context_cost(n: int, layers: int = LAYERS, kv_heads: int = KV_HEADS, head_dim: int = HEAD_DIM, bits: int = 16) -> dict:
1370    """The two costs of an n-token context: query-key pairs scored, and KV-cache bytes held.
1371
1372    Pairs are counted per head per layer and without the causal halving, as
1373    `primer.ml.attention.attention_cost` does: the point is the n² growth.
1374    """
1375    return dict(n=n, pairs_per_head_per_layer=n * n, kv_cache_bytes=kv_cache_bytes(n, layers, kv_heads, head_dim, bits))
1376
1377
1378# ---------------------------------------------------------------------------
1379# 2. Sliding-window attention
1380# ---------------------------------------------------------------------------
1381
1382
1383def sliding_window_mask(n: int, w: int) -> np.ndarray:
1384    """Boolean (n, n) mask, True where query i may read key j: the last w tokens, itself included.
1385
1386    For n = 6 and w = 3:
1387
1388    ```text
1389    [[1 0 0 0 0 0]
1390     [1 1 0 0 0 0]
1391     [1 1 1 0 0 0]
1392     [0 1 1 1 0 0]
1393     [0 0 1 1 1 0]
1394     [0 0 0 1 1 1]]
1395    ```
1396    """
1397    back = np.arange(n)[:, None] - np.arange(n)[None, :]  # how far back key j is from query i
1398    return (back >= 0) & (back < w)  # never the future, never further back than w − 1
1399
1400
1401def pairs_computed(mask: np.ndarray) -> int:
1402    """How many query-key scores a mask asks for: the work attention actually does."""
1403    return int(np.count_nonzero(mask))
1404
1405
1406def receptive_field(n: int, w: int, layers: int) -> np.ndarray:
1407    """Boolean (n, n): True where token j can influence token i's output after `layers` windowed layers.
1408
1409    One layer moves information along the mask's edges; stacking layers is
1410    following edges several times, which is repeated matrix multiplication
1411    of the 0/1 mask (clipped back to 0/1 so the counts stay small).
1412    """
1413    step = sliding_window_mask(n, w).astype(np.int64)
1414    heard = np.eye(n, dtype=np.int64)  # before any layer, each token knows only itself
1415    for _ in range(layers):
1416        heard = np.minimum(step @ heard, 1)
1417    return heard.astype(bool)
1418
1419
1420def reach(w: int, layers: int) -> int:
1421    """How many positions back information can travel: each layer adds w − 1."""
1422    return layers * (w - 1)
1423
1424
1425def sliding_window_decode(Q: np.ndarray, K: np.ndarray, V: np.ndarray, w: int) -> tuple[np.ndarray, int]:
1426    """Generate token by token with a rolling KV buffer that holds at most w entries.
1427
1428    Returns (outputs (n, d_v), the largest the buffer ever got). The outputs
1429    equal masked attention with `sliding_window_mask`; the memory never grows past w.
1430    """
1431    keys: deque = deque(maxlen=w)  # appending the (w+1)-th entry silently drops the oldest
1432    values: deque = deque(maxlen=w)
1433    outputs, largest = [], 0
1434    for q, k, v in zip(Q, K, V):
1435        keys.append(k)
1436        values.append(v)
1437        largest = max(largest, len(keys))
1438        weights = softmax(np.array(keys) @ q / np.sqrt(len(q)))  # one row of attention, over the buffer only
1439        outputs.append(weights @ np.array(values))
1440    return np.array(outputs), largest
1441
1442
1443# ---------------------------------------------------------------------------
1444# 3. Sparse patterns: local + global, strided
1445# ---------------------------------------------------------------------------
1446
1447
1448def global_local_mask(n: int, w: int, global_tokens: tuple[int, ...] = (0,)) -> np.ndarray:
1449    """A sliding window plus a few global tokens that everyone reads and that read everyone (causally)."""
1450    mask = sliding_window_mask(n, w)
1451    g = list(global_tokens)
1452    mask[:, g] = True  # every token may read the global tokens...
1453    mask[g, :] = True  # ...and the global tokens may read every token...
1454    return mask & causal_mask(n)  # ...but nobody reads the future
1455
1456
1457def strided_mask(n: int, stride: int) -> np.ndarray:
1458    """The Sparse Transformer's strided pattern: the last `stride` tokens, plus every stride-th token before them."""
1459    back = np.arange(n)[:, None] - np.arange(n)[None, :]
1460    return (back >= 0) & ((back < stride) | (back % stride == 0))
1461
1462
1463# ---------------------------------------------------------------------------
1464# 4. Linear attention
1465# ---------------------------------------------------------------------------
1466
1467
1468def feature_map(x: np.ndarray) -> np.ndarray:
1469    """φ(x) = elu(x) + 1: x + 1 for positive x, e^x otherwise. Always positive, so weights never go negative."""
1470    x = np.asarray(x, dtype=float)
1471    return np.where(x > 0, x + 1.0, np.exp(np.minimum(x, 0.0)))  # the minimum keeps exp from overflowing on the unused branch
1472
1473
1474def linear_attention_weights(Q: np.ndarray, K: np.ndarray) -> np.ndarray:
1475    """The (n, n) causal weights linear attention implies: φ(q_i)·φ(k_j), each row divided by its total."""
1476    scores = feature_map(Q) @ feature_map(K).T  # positive "similarities", no exponential
1477    scores = np.where(causal_mask(len(Q)), scores, 0.0)  # the future contributes nothing
1478    return scores / scores.sum(axis=1, keepdims=True)
1479
1480
1481def linear_attention_quadratic(Q: np.ndarray, K: np.ndarray, V: np.ndarray) -> np.ndarray:
1482    """Linear attention written the slow way: build all n × n weights, then blend the values."""
1483    return linear_attention_weights(Q, K) @ V
1484
1485
1486def linear_attention_recurrent(Q: np.ndarray, K: np.ndarray, V: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
1487    """Linear attention the fast way: two running sums, updated once per token.
1488
1489    S (d_k, d_v) accumulates φ(k_j) v_jᵀ and z (d_k,) accumulates φ(k_j).
1490    Token i reads φ(q_i)ᵀ S / φ(q_i)ᵀ z. Returns (outputs, final S, final z):
1491    the state is the same size at token 3 and at token 3 million.
1492    """
1493    S = np.zeros((K.shape[1], V.shape[1]))
1494    z = np.zeros(K.shape[1])
1495    out = np.zeros((len(Q), V.shape[1]))
1496    for i, (q, k, v) in enumerate(zip(Q, K, V)):
1497        phi_k = feature_map(k)
1498        S = S + np.outer(phi_k, v)  # pour this token's value into the pot, flavoured by its key
1499        z = z + phi_k  # and remember how much flavour went in, to normalise later
1500        phi_q = feature_map(q)
1501        out[i] = (phi_q @ S) / (phi_q @ z)
1502    return out, S, z
1503
1504
1505# ---------------------------------------------------------------------------
1506# 5. State-space models
1507# ---------------------------------------------------------------------------
1508
1509
1510def _ssm_params(A, B, C) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
1511    # Accept plain numbers for the one-number worked example, matrices for the real thing.
1512    return np.atleast_2d(np.asarray(A, float)), np.atleast_1d(np.asarray(B, float)), np.atleast_1d(np.asarray(C, float))
1513
1514
1515def ssm_recurrent(x, A, B, C) -> np.ndarray:
1516    """h_t = A h_{t-1} + B x_t, y_t = C · h_t, one step at a time (how an SSM runs during generation).
1517
1518    x: (T,) one input channel. A: (N, N) or a number; B, C: (N,) or numbers.
1519    """
1520    A, B, C = _ssm_params(A, B, C)
1521    h = np.zeros(len(B))  # the state: N numbers, whatever the sequence length
1522    y = []
1523    for x_t in np.asarray(x, float):
1524        h = A @ h + B * x_t
1525        y.append(float(C @ h))
1526    return np.array(y)
1527
1528
1529def ssm_kernel(A, B, C, length: int) -> np.ndarray:
1530    """The convolution kernel (C B, C A B, C A² B, ...): how much an input k steps ago still counts."""
1531    A, B, C = _ssm_params(A, B, C)
1532    kernel, A_power_B = [], B.copy()
1533    for _ in range(length):
1534        kernel.append(float(C @ A_power_B))
1535        A_power_B = A @ A_power_B
1536    return np.array(kernel)
1537
1538
1539def ssm_convolution(x, A, B, C) -> np.ndarray:
1540    """The same outputs as `ssm_recurrent`, computed as one convolution (how a fixed SSM trains in parallel)."""
1541    x = np.asarray(x, float)
1542    kernel = ssm_kernel(A, B, C, len(x))
1543    # y_t = Σ_k kernel_k · x_{t−k}; np.convolve does exactly that sum, and we keep the first T outputs.
1544    return np.convolve(x, kernel)[: len(x)]
1545
1546
1547def discretize(delta: float, a: float = -1.0, b: float = 1.0) -> tuple[float, float]:
1548    """Turn a step size Δ into this step's keep and write factors (zero-order hold, one channel).
1549
1550    Ā = e^(Δa), B̄ = (e^(Δa) − 1) / a · b. With a = −1 and b = 1 they add up to 1:
1551    a big Δ overwrites the state, a tiny Δ leaves it alone.
1552    """
1553    A_bar = float(np.exp(delta * a))
1554    return A_bar, float((A_bar - 1.0) / a * b)
1555
1556
1557def softplus(x):
1558    """log(1 + e^x): a smooth ramp that is always positive, so a step size can never go negative."""
1559    return np.logaddexp(0.0, x)  # log(e^0 + e^x) without overflow
1560
1561
1562def step_sizes(markers, sharpness: float = 10.0, bias: float = -5.0) -> np.ndarray:
1563    """Selectivity in miniature: Δ_t = softplus(sharpness · marker_t + bias), read off each token itself.
1564
1565    Marked tokens get Δ ≈ 5 (write), unmarked ones Δ ≈ 0.007 (ignore).
1566    Mamba learns this map; here it is set by hand so the numbers are checkable.
1567    """
1568    return softplus(sharpness * np.asarray(markers, float) + bias)
1569
1570
1571def selective_ssm(x, delta, a: float = -1.0, b: float = 1.0, c: float = 1.0) -> np.ndarray:
1572    """One-channel selective SSM: every step discretizes with its own Δ_t, so every step has its own Ā_t and B̄_t."""
1573    h, ys = 0.0, []
1574    for x_t, d_t in zip(np.asarray(x, float), np.asarray(delta, float)):
1575        A_bar, B_bar = discretize(d_t, a, b)
1576        h = A_bar * h + B_bar * x_t
1577        ys.append(c * h)
1578    return np.array(ys)
1579
1580
1581# The recall task: one marked token (the 7) among noise, then nine more noise tokens.
1582RECALL_VALUES = np.array([0.3, -0.5, 7.0, 0.8, -0.2, 0.6, -0.9, 0.4, 0.1, -0.3, 0.5, -0.6])
1583RECALL_MARKERS = np.array([0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0])
1584
1585
1586def selective_recall() -> float:
1587    """The selective SSM's state after the recall task: how much of the marked 7 survived."""
1588    return float(selective_ssm(RECALL_VALUES, step_sizes(RECALL_MARKERS))[-1])
1589
1590
1591def time_invariant_recall(delta: float) -> float:
1592    """The same task with one fixed Δ for every token (a time-invariant SSM, like S4)."""
1593    return float(selective_ssm(RECALL_VALUES, np.full(len(RECALL_VALUES), delta))[-1])
1594
1595
1596def parallel_scan(a, b) -> tuple[np.ndarray, int]:
1597    """Every state of h_t = a_t h_{t-1} + b_t (h before the start = 0) in ceil(log2 T) rounds.
1598
1599    Each step is the map h -> a h + b. Two steps in a row are again such a map:
1600    first (a1, b1) then (a2, b2) is (a1·a2, a2·b1 + b2). In round r every
1601    position combines with the one 2^r places earlier, all positions at once,
1602    so on parallel hardware the loop over T steps becomes log2 T rounds.
1603    Returns (states (T,), rounds).
1604    """
1605    a, b = np.array(a, float), np.array(b, float)
1606    shift, rounds = 1, 0
1607    while shift < len(a):
1608        # The step `shift` places earlier; before the start it is the do-nothing map (1, 0).
1609        a_prev = np.concatenate([np.ones(shift), a[:-shift]])
1610        b_prev = np.concatenate([np.zeros(shift), b[:-shift]])
1611        a, b = a_prev * a, a * b_prev + b  # both right-hand sides use the old a
1612        shift, rounds = shift * 2, rounds + 1
1613    return b, rounds
1614
1615
1616def ssm_state_bytes(channels: int, state_size: int, bits: int = 16) -> int:
1617    """Memory an SSM layer carries between tokens: channels × state numbers, at any context length."""
1618    return channels * state_size * bits // 8
1619
1620
1621def hybrid_cache_bytes(
1622    context: int, layers: int = LAYERS, attention_layers: int = 4, kv_heads: int = KV_HEADS, head_dim: int = HEAD_DIM,
1623    channels: int = 4096, state_size: int = 16, bits: int = 16,
1624) -> int:
1625    """A hybrid stack: attention layers keep a growing KV cache, SSM layers a fixed state."""
1626    attention = attention_layers * context * kv_cache_bytes_per_token(1, kv_heads, head_dim, bits)
1627    ssm = (layers - attention_layers) * ssm_state_bytes(channels, state_size, bits)
1628    return attention + ssm
1629
1630
1631# ---------------------------------------------------------------------------
1632# 6. Compressing the KV cache: latents and quantization
1633# ---------------------------------------------------------------------------
1634
1635
1636def latent_kv_bytes_per_token(layers: int, latent_dim: int, rope_dim: int = 0, bits: int = 16) -> int:
1637    """Latent KV cache per token: one latent (plus a small position-carrying key) per layer."""
1638    return layers * (latent_dim + rope_dim) * bits // 8
1639
1640
1641# The worked example: a 4-number token, a 2-number latent, 2 heads of width 2.
1642WORKED_X = np.array([1.0, 0.0, 2.0, 1.0])
1643WORKED_W_DOWN = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 0.0], [0.0, 1.0]])
1644WORKED_W_UK = np.array([[1.0, 0.0, 1.0, 0.0], [0.0, 1.0, 0.0, 1.0]])
1645
1646
1647def latent_kv_worked_example() -> dict[str, np.ndarray]:
1648    """x = (1, 0, 2, 1) is cached as the latent (3, 1), which expands to the key (3, 1, 3, 1) when needed."""
1649    latent = WORKED_X @ WORKED_W_DOWN
1650    return dict(latent=latent, key=latent @ WORKED_W_UK)
1651
1652
1653class LatentKVAttention:
1654    """Multi-head attention that caches one small latent per token instead of every head's keys and values.
1655
1656    The latent c = x W_down (d_latent numbers) is all that is stored. Keys
1657    and values are rebuilt from it: K = C W_uk and V = C W_uv, for all heads.
1658    Weights are random: we are studying the mechanics, not training.
1659    """
1660
1661    def __init__(self, d_model: int, n_heads: int, d_head: int, d_latent: int, seed: int = 0):
1662        rng = np.random.default_rng(seed)
1663        self.n_heads, self.d_head, self.d_latent = n_heads, d_head, d_latent
1664        s, s_latent = 1 / np.sqrt(d_model), 1 / np.sqrt(d_latent)  # keep activations near unit size
1665        self.W_q = rng.normal(0, s, (d_model, n_heads * d_head))
1666        self.W_down = rng.normal(0, s, (d_model, d_latent))  # squeeze: what gets cached
1667        self.W_uk = rng.normal(0, s_latent, (d_latent, n_heads * d_head))  # expand to every head's keys
1668        self.W_uv = rng.normal(0, s_latent, (d_latent, n_heads * d_head))  # and to every head's values
1669
1670    def _heads(self, M: np.ndarray) -> np.ndarray:
1671        # (n, H·d_head) -> (H, n, d_head): one slab per head.
1672        return M.reshape(M.shape[0], self.n_heads, self.d_head).transpose(1, 0, 2)
1673
1674    def _merge(self, M: np.ndarray) -> np.ndarray:
1675        # (H, n, d_head) -> (n, H·d_head): heads side by side again.
1676        return M.transpose(1, 0, 2).reshape(M.shape[1], self.n_heads * self.d_head)
1677
1678    def compress(self, X: np.ndarray) -> np.ndarray:
1679        """(n, d_model) -> (n, d_latent): the whole KV cache for these tokens."""
1680        return X @ self.W_down
1681
1682    def attend(self, X: np.ndarray, C: np.ndarray) -> np.ndarray:
1683        """Causal attention the direct way: expand the cached latents into keys and values, then attend."""
1684        Q, K, V = self._heads(X @ self.W_q), self._heads(C @ self.W_uk), self._heads(C @ self.W_uv)
1685        out, _ = scaled_dot_product_attention(Q, K, V, mask=causal_mask(len(X)))
1686        return self._merge(out)
1687
1688    def attend_absorbed(self, X: np.ndarray, C: np.ndarray) -> np.ndarray:
1689        """The same result without ever building K or V: fold W_uk into the query and W_uv into the output.
1690
1691        q · (c W_uk) = (q W_ukᵀ) · c, and Σ w_j (c_j W_uv) = (Σ w_j c_j) W_uv.
1692        """
1693        Q = self._heads(X @ self.W_q)  # (H, n, d_head)
1694        W_uk = self.W_uk.reshape(self.d_latent, self.n_heads, self.d_head).transpose(1, 0, 2)  # (H, d_latent, d_head)
1695        W_uv = self.W_uv.reshape(self.d_latent, self.n_heads, self.d_head).transpose(1, 0, 2)
1696        Q_latent = Q @ W_uk.transpose(0, 2, 1)  # (H, n, d_latent): each query translated into latent space
1697        scores = Q_latent @ C.T / np.sqrt(self.d_head)  # scores against the latents directly
1698        scores = np.where(causal_mask(len(X)), scores, -np.inf)
1699        mixed = softmax(scores) @ C  # (H, n, d_latent): blend the latents first...
1700        return self._merge(mixed @ W_uv)  # ...then expand once per query, not once per cached token
1701
1702    def cached_numbers_per_token(self) -> int:
1703        return self.d_latent
1704
1705    def full_numbers_per_token(self) -> int:
1706        """What ordinary multi-head attention would cache: a key and a value for every head."""
1707        return 2 * self.n_heads * self.d_head
1708
1709
1710def quantize_kv(M: np.ndarray, bits: int) -> tuple[np.ndarray, np.ndarray]:
1711    """Quantize cached keys or values with one scale per row (per token, per head): codes and scales."""
1712    return quantize(M, bits, per_channel=True)  # `primer.ml.inference.quantize`: per_channel means one scale per row
1713
1714
1715def quantized_kv_bytes_per_token(layers: int, kv_heads: int, head_dim: int, bits: int, scale_bits: int = 16) -> float:
1716    """2 (key and value) × layers × KV heads × (head_dim codes of `bits` each + one scale)."""
1717    return 2 * layers * kv_heads * (head_dim * bits / 8 + scale_bits / 8)
1718
1719
1720def quantized_cache_error(bits: int, n: int = 64, d: int = 64, seed: int = 0) -> float:
1721    """Relative change in causal attention output when K and V are stored at `bits` bits instead of full precision."""
1722    rng = np.random.default_rng(seed)
1723    Q, K, V = (rng.standard_normal((n, d)) for _ in range(3))
1724    exact, _ = scaled_dot_product_attention(Q, K, V, mask=causal_mask(n))
1725    K_hat, V_hat = dequantize(*quantize_kv(K, bits)), dequantize(*quantize_kv(V, bits))
1726    approx, _ = scaled_dot_product_attention(Q, K_hat, V_hat, mask=causal_mask(n))
1727    return float(np.linalg.norm(approx - exact) / np.linalg.norm(exact))
1728
1729
1730def cache_bytes_by_method(context: int) -> dict[str, float]:
1731    """Memory to hold one conversation's past, for each design, on the 32-layer running example."""
1732    return {
1733        "multi-head (32 KV heads)": kv_cache_bytes(context, LAYERS, 32, HEAD_DIM),
1734        "grouped-query (8 KV heads)": kv_cache_bytes(context, LAYERS, KV_HEADS, HEAD_DIM),
1735        "grouped-query, 4-bit": context * quantized_kv_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM, bits=4),
1736        "latent (576 per layer)": context * latent_kv_bytes_per_token(LAYERS, 512, 64),
1737        "sliding window 4k": kv_cache_bytes(min(context, 4096), LAYERS, KV_HEADS, HEAD_DIM),
1738        "hybrid, 1 attention in 8": hybrid_cache_bytes(context),
1739        "state-space": LAYERS * ssm_state_bytes(4096, 16),
1740    }
1741
1742
1743# ---------------------------------------------------------------------------
1744# 7. Figures (rendered into the HTML docs by `make figures`)
1745# ---------------------------------------------------------------------------
1746
1747
1748def figures() -> dict:
1749    """Plot this lesson's data. matplotlib is imported here, and only here,
1750    so the lesson itself needs nothing beyond NumPy."""
1751    import matplotlib
1752
1753    matplotlib.use("Agg")
1754    import matplotlib.pyplot as plt
1755
1756    BLUE, RED, ORANGE, GREEN, PURPLE, MUTED = "#2563eb", "#dc2626", "#d97706", "#059669", "#7c3aed", "#9ca3af"
1757    figs = {}
1758
1759    # --- 1. The two bills at 8k, 128k and 1M tokens ------------------------------
1760    contexts = (8_192, 131_072, 1_048_576)
1761    labels = ("8k", "128k", "1M")
1762    costs = [context_cost(n) for n in contexts]
1763    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4))
1764    a1.bar(labels, [c["pairs_per_head_per_layer"] for c in costs], color=RED)
1765    a1.set_yscale("log")
1766    a1.set_ylabel("pairs per head per layer")
1767    a1.set_title("Compute: pairs scored grow with n²")
1768    gb = [c["kv_cache_bytes"] / 1e9 for c in costs]
1769    a2.bar(labels, gb, color=BLUE)
1770    for i, g in enumerate(gb):
1771        a2.text(i, g + 2, f"{g:.1f} GB", ha="center")
1772    a2.axhline(80, color=MUTED, ls="--")
1773    a2.text(-0.4, 84, "one 80 GB GPU", color="#4b5563")
1774    a2.set_ylabel("KV cache, GB")
1775    a2.set_ylim(0, 160)
1776    a2.set_title("Memory: the KV cache grows with n")
1777    for a in (a1, a2):
1778        a.set_xlabel("context length")
1779    fig.tight_layout()
1780    figs["context_cost"] = fig
1781
1782    # --- 2. Sliding window: one layer vs. what three layers can reach -----------
1783    fig, (a1, a2) = plt.subplots(1, 2, figsize=(8, 3.8))
1784    a1.imshow(sliding_window_mask(16, 4), cmap="Blues", vmin=0, vmax=1.3)
1785    a1.set_title("One layer, window w = 4")
1786    a2.imshow(receptive_field(16, 4, layers=3), cmap="Blues", vmin=0, vmax=1.3)
1787    a2.set_title("Reach after 3 layers: 3·(4 − 1) = 9 back")
1788    for a in (a1, a2):
1789        a.set_xlabel("key position (being read)")
1790        a.set_ylabel("query position (reading)")
1791        a.set_xticks([0, 5, 10, 15])
1792        a.set_yticks([0, 5, 10, 15])
1793        a.grid(False)
1794    fig.tight_layout()
1795    figs["window_reach"] = fig
1796
1797    # --- 3. Sparse masks and the pairs they score --------------------------------
1798    masks = [
1799        ("full causal", causal_mask(16)),
1800        ("window w = 4", sliding_window_mask(16, 4)),
1801        ("window + global token 0", global_local_mask(16, 4, (0,))),
1802        ("strided, stride 4", strided_mask(16, 4)),
1803    ]
1804    fig, axes = plt.subplots(1, 4, figsize=(12, 3.4))
1805    for a, (name, m) in zip(axes, masks):
1806        a.imshow(m, cmap="Blues", vmin=0, vmax=1.3)
1807        a.set_title(f"{name}\n{pairs_computed(m)} pairs")
1808        a.set_xticks([0, 5, 10, 15])
1809        a.set_yticks([0, 5, 10, 15])
1810        a.set_xlabel("key")
1811        a.grid(False)
1812    axes[0].set_ylabel("query")
1813    fig.tight_layout()
1814    figs["sparse_masks"] = fig
1815
1816    # --- 4. Softmax vs. linear attention weights on the same inputs -------------
1817    rng = np.random.default_rng(11)
1818    # Queries twice the usual size: decisive, as trained queries often are. Softmax sharpens; linear barely moves.
1819    Q, K = 2 * rng.standard_normal((10, 8)), rng.standard_normal((10, 8))
1820    _, soft = scaled_dot_product_attention(Q, K, K, mask=causal_mask(10))
1821    lin = linear_attention_weights(Q, K)
1822    fig, (a1, a2) = plt.subplots(1, 2, figsize=(8.4, 3.8))
1823    for a, w, title in ((a1, soft, "softmax attention"), (a2, lin, "linear attention, φ = elu + 1")):
1824        im = a.imshow(w, cmap="Blues", vmin=0, vmax=1)
1825        a.set_title(title)
1826        a.set_xlabel("key")
1827        a.set_ylabel("query")
1828        a.grid(False)
1829    fig.colorbar(im, ax=[a1, a2], fraction=0.03, label="attention weight")
1830    figs["linear_vs_softmax"] = fig
1831
1832    # --- 5. SSM kernels: A sets how long an input keeps counting ----------------
1833    steps = np.arange(151)
1834    fig, ax = plt.subplots(figsize=(6, 3.4))
1835    for A, color in ((0.5, RED), (0.9, ORANGE), (0.99, BLUE)):
1836        ax.plot(steps, ssm_kernel(A, 1.0, 1.0, len(steps)), color=color, label=f"A = {A}")
1837    ax.set_xlabel("steps since the input arrived (k)")
1838    ax.set_ylabel("how much it still counts, C·A^k·B")
1839    ax.set_title("The kernel: how fast a fixed SSM forgets")
1840    ax.legend(frameon=False)
1841    fig.tight_layout()
1842    figs["ssm_kernel"] = fig
1843
1844    # --- 6. Selective recall ------------------------------------------------------
1845    grid = np.geomspace(1e-3, 10, 60)
1846    best = grid[int(np.argmax([time_invariant_recall(d) for d in grid]))]
1847    t = np.arange(len(RECALL_VALUES))
1848    fig, ax = plt.subplots(figsize=(7, 3.6))
1849    ax.bar(t, RECALL_VALUES, color=MUTED, label="input (marked 7 at position 2)")
1850    ax.plot(t, selective_ssm(RECALL_VALUES, step_sizes(RECALL_MARKERS)), "o-", color=BLUE, label="selective: Δ set by each token")
1851    ax.plot(t, selective_ssm(RECALL_VALUES, np.full(len(t), best)), "o-", color=RED, label=f"best fixed Δ = {best:.2f}")
1852    ax.plot(t, selective_ssm(RECALL_VALUES, np.full(len(t), 5.0)), "o-", color=ORANGE, label="large fixed Δ = 5")
1853    ax.set_xlabel("position in the sequence")
1854    ax.set_ylabel("value / state")
1855    ax.set_title("Remembering one marked token through noise")
1856    ax.legend(frameon=False, fontsize=8, loc="center right")
1857    fig.tight_layout()
1858    figs["selective_recall"] = fig
1859
1860    # --- 7. Cache per token, by design ------------------------------------------
1861    designs = [
1862        ("multi-head, 32 KV heads", kv_cache_bytes_per_token(LAYERS, 32, HEAD_DIM)),
1863        ("grouped-query, 8 KV heads", kv_cache_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM)),
1864        ("latent, 512 + 64 numbers", latent_kv_bytes_per_token(LAYERS, 512, 64)),
1865        ("grouped-query 8, 4-bit", quantized_kv_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM, bits=4)),
1866        ("multi-query, 1 KV head", kv_cache_bytes_per_token(LAYERS, 1, HEAD_DIM)),
1867    ]
1868    fig, ax = plt.subplots(figsize=(7, 3.2))
1869    kb = [b / 1024 for _, b in designs]
1870    ax.barh([d for d, _ in designs][::-1], kb[::-1], color=[GREEN, ORANGE, PURPLE, BLUE, RED])
1871    for i, v in enumerate(kb[::-1]):
1872        ax.text(v + 5, i, f"{v:.0f} KB", va="center")
1873    ax.set_xlabel("KV cache per token, all 32 layers (KB)")
1874    ax.set_xlim(0, 600)
1875    ax.set_title("Shrinking the cache per token")
1876    fig.tight_layout()
1877    figs["kv_per_token"] = fig
1878
1879    # --- 8. Memory vs. context, every design ------------------------------------
1880    ns = np.geomspace(1_000, 1_048_576, 80).astype(int)
1881    rows = [cache_bytes_by_method(int(n)) for n in ns]
1882    fig, ax = plt.subplots(figsize=(7.5, 4.2))
1883    for name in rows[0]:
1884        ax.loglog(ns, [r[name] / 1e9 for r in rows], label=name)
1885    ax.axhline(80, color=MUTED, ls="--")
1886    ax.text(1_200, 95, "80 GB GPU", color="#4b5563")
1887    ax.set_xlabel("context length n (tokens)")
1888    ax.set_ylabel("cache for one conversation (GB)")
1889    ax.set_title("What each design must keep, as context grows")
1890    ax.legend(frameon=False, fontsize=8, loc="lower right")
1891    fig.tight_layout()
1892    figs["memory_vs_context"] = fig
1893
1894    return figs
1895
1896
1897# ---------------------------------------------------------------------------
1898# 8. Narrated walkthrough
1899# ---------------------------------------------------------------------------
1900
1901
1902def demo() -> None:
1903    banner("1. Why long context is expensive")
1904    table(
1905        ["context", "pairs per head per layer", "KV cache (GB)"],
1906        [(f"{c['n']:,}", f"{c['pairs_per_head_per_layer']:,}", c["kv_cache_bytes"] / 1e9)
1907         for c in (context_cost(n) for n in (8_192, 131_072, 1_048_576))],
1908        floatfmt=".1f",
1909    )
1910    takeaway("Pairs grow with n², the cache with n. At 1M tokens the cache alone overflows an 80 GB GPU.")
1911
1912    banner("2. Sliding-window attention")
1913    matrix("window w = 3 over 8 tokens (1 = score computed)", sliding_window_mask(8, 3).astype(int))
1914    say(
1915        f"""
1916        {pairs_computed(sliding_window_mask(8, 3))} pairs instead of {pairs_computed(causal_mask(8))}.
1917        After three layers, token 7 can hear token 1:
1918        {bool(receptive_field(8, 3, 3)[7, 1])}; token 0: {bool(receptive_field(8, 3, 3)[7, 0])}.
1919        Mistral 7B's 32 layers of w = 4,096 reach {reach(4096, 32):,} tokens back.
1920        """
1921    )
1922    rng = np.random.default_rng(0)
1923    Q, K, V = (rng.standard_normal((50, 8)) for _ in range(3))
1924    streamed, largest = sliding_window_decode(Q, K, V, w=4)
1925    masked, _ = scaled_dot_product_attention(Q, K, V, mask=sliding_window_mask(50, 4))
1926    say(f"Decoding 50 tokens with a rolling buffer: largest buffer {largest}, same outputs: {np.allclose(streamed, masked)}.")
1927    takeaway("A window caps both bills; depth still carries information L·(w − 1) positions back.")
1928
1929    banner("3. Sparse patterns: count the pairs")
1930    table(
1931        ["pattern (16 tokens)", "pairs scored"],
1932        [
1933            ("full causal", pairs_computed(causal_mask(16))),
1934            ("window w = 4", pairs_computed(sliding_window_mask(16, 4))),
1935            ("window + global token 0", pairs_computed(global_local_mask(16, 4, (0,)))),
1936            ("strided, stride 4", pairs_computed(strided_mask(16, 4))),
1937        ],
1938    )
1939    table(
1940        ["tokens", "full causal", "window 64", "ratio"],
1941        [(n, pairs_computed(causal_mask(n)), pairs_computed(sliding_window_mask(n, 64)),
1942          pairs_computed(causal_mask(n)) / pairs_computed(sliding_window_mask(n, 64))) for n in (256, 1024, 4096)],
1943        floatfmt=".1f",
1944    )
1945    takeaway("With a fixed window the pairs grow with n, not n²; global tokens keep everyone two hops apart.")
1946
1947    banner("4. Linear attention: the pot")
1948    Qw = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 0.0]])
1949    Kw = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]])
1950    Vw = np.array([[2.0], [4.0], [6.0]])
1951    out, S, z = linear_attention_recurrent(Qw, Kw, Vw)
1952    soft, _ = scaled_dot_product_attention(Qw, Kw, Vw, mask=causal_mask(3), scale=False)
1953    say(
1954        f"""
1955        Keys (1,0), (0,1), (1,1); values 2, 4, 6; the third query is (1,0).
1956        The pot S = {S[:, 0].tolist()}, the total z = {z.tolist()}, so the output is
1957        {out[2, 0]:.3f} (= 62/15). The same inputs through softmax attention give {soft[2, 0]:.3f}.
1958        """
1959    )
1960    Q, K, V = (rng.standard_normal((200, 8)) for _ in range(3))
1961    fast, S_big, z_big = linear_attention_recurrent(Q, K, V)
1962    same = np.allclose(fast, linear_attention_quadratic(Q, K, V))
1963    say(f"On 200 random tokens the running sum equals the n × n form: {same}. The state is {S_big.shape} plus {z_big.shape}, at any length.")
1964    takeaway("Swap e^(q·k) for φ(q)·φ(k) and attention becomes a running sum: O(n), fixed state, blurrier focus.")
1965
1966    banner("5. State-space models")
1967    say(
1968        f"""
1969        A = 0.5, inputs (1, 0, 0, 2). Recurrence: {ssm_recurrent([1, 0, 0, 2], 0.5, 1.0, 1.0).tolist()}.
1970        Kernel: {ssm_kernel(0.5, 1.0, 1.0, 4).tolist()}. Convolution: {ssm_convolution([1, 0, 0, 2], 0.5, 1.0, 1.0).tolist()}.
1971        """
1972    )
1973    table(
1974        ["step size Δ", "keep Ā", "write B̄"],
1975        [(d, *discretize(d)) for d in (0.01, 0.1, 1.0, 5.0)],
1976    )
1977    grid = np.geomspace(1e-3, 10, 60)
1978    say(
1979        f"""
1980        Recall task (a marked 7, then nine noise tokens): the selective SSM ends at
1981        {selective_recall():.2f}; the best fixed step ends at
1982        {max(time_invariant_recall(d) for d in grid):.2f}.
1983        """
1984    )
1985    states, rounds = parallel_scan(np.full(1024, 0.99), rng.standard_normal(1024))
1986    say(f"A parallel scan over 1,024 steps finishes in {rounds} rounds.")
1987    table(
1988        ["context", "all-attention cache (GB)", "hybrid 4 of 32 (GB)", "pure SSM state (MB)"],
1989        [(f"{n:,}", context_cost(n)["kv_cache_bytes"] / 1e9, hybrid_cache_bytes(n) / 1e9, LAYERS * ssm_state_bytes(4096, 16) / 1e6)
1990         for n in (8_192, 131_072, 1_048_576)],
1991        floatfmt=".2f",
1992    )
1993    takeaway("A selective SSM decides per token what to keep; its memory never grows, and a few attention layers restore exact recall.")
1994
1995    banner("6. Compressing the KV cache")
1996    worked = latent_kv_worked_example()
1997    say(f"Latent example: x = (1, 0, 2, 1) is cached as {worked['latent'].tolist()} and expands to the key {worked['key'].tolist()}.")
1998    layer = LatentKVAttention(d_model=64, n_heads=8, d_head=8, d_latent=16)
1999    X = rng.standard_normal((12, 64))
2000    C = layer.compress(X)
2001    say(
2002        f"""
2003        A latent layer with 8 heads of width 8 caches {layer.cached_numbers_per_token()} numbers per token
2004        instead of {layer.full_numbers_per_token()}. Expanded and absorbed attention agree:
2005        {np.allclose(layer.attend(X, C), layer.attend_absorbed(X, C))}.
2006        """
2007    )
2008    table(
2009        ["cache precision", "output error", "bytes per token (8B example)"],
2010        [(f"{b}-bit", quantized_cache_error(b), f"{quantized_kv_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM, b, scale_bits=0 if b == 16 else 16):,.0f}")
2011         for b in (16, 8, 4)],
2012    )
2013    table(
2014        ["design", "cache at 1M tokens (GB)"],
2015        [(name, b / 1e9) for name, b in cache_bytes_by_method(1_048_576).items()],
2016        floatfmt=".3f",
2017    )
2018    takeaway("Share heads, squeeze into a latent, cut bits, or bound the window: each moves the memory line down or flattens it.")
2019
2020
2021if __name__ == "__main__":
2022    demo()
Level 3: the code, function by function.
def context_cost( n: int, layers: int = 32, kv_heads: int = 8, head_dim: int = 128, bits: int = 16) -> dict: on GitHub
1370def context_cost(n: int, layers: int = LAYERS, kv_heads: int = KV_HEADS, head_dim: int = HEAD_DIM, bits: int = 16) -> dict:
1371    """The two costs of an n-token context: query-key pairs scored, and KV-cache bytes held.
1372
1373    Pairs are counted per head per layer and without the causal halving, as
1374    `primer.ml.attention.attention_cost` does: the point is the n² growth.
1375    """
1376    return dict(n=n, pairs_per_head_per_layer=n * n, kv_cache_bytes=kv_cache_bytes(n, layers, kv_heads, head_dim, bits))

The two costs of an n-token context: query-key pairs scored, and KV-cache bytes held.

Pairs are counted per head per layer and without the causal halving, as primer.ml.attention.attention_cost does: the point is the n² growth.

def sliding_window_mask(n: int, w: int) -> numpy.ndarray: on GitHub
1384def sliding_window_mask(n: int, w: int) -> np.ndarray:
1385    """Boolean (n, n) mask, True where query i may read key j: the last w tokens, itself included.
1386
1387    For n = 6 and w = 3:
1388
1389    ```text
1390    [[1 0 0 0 0 0]
1391     [1 1 0 0 0 0]
1392     [1 1 1 0 0 0]
1393     [0 1 1 1 0 0]
1394     [0 0 1 1 1 0]
1395     [0 0 0 1 1 1]]
1396    ```
1397    """
1398    back = np.arange(n)[:, None] - np.arange(n)[None, :]  # how far back key j is from query i
1399    return (back >= 0) & (back < w)  # never the future, never further back than w − 1

Boolean (n, n) mask, True where query i may read key j: the last w tokens, itself included.

For n = 6 and w = 3:

[[1 0 0 0 0 0]
 [1 1 0 0 0 0]
 [1 1 1 0 0 0]
 [0 1 1 1 0 0]
 [0 0 1 1 1 0]
 [0 0 0 1 1 1]]
def pairs_computed(mask: numpy.ndarray) -> int: on GitHub
1402def pairs_computed(mask: np.ndarray) -> int:
1403    """How many query-key scores a mask asks for: the work attention actually does."""
1404    return int(np.count_nonzero(mask))

How many query-key scores a mask asks for: the work attention actually does.

def receptive_field(n: int, w: int, layers: int) -> numpy.ndarray: on GitHub
1407def receptive_field(n: int, w: int, layers: int) -> np.ndarray:
1408    """Boolean (n, n): True where token j can influence token i's output after `layers` windowed layers.
1409
1410    One layer moves information along the mask's edges; stacking layers is
1411    following edges several times, which is repeated matrix multiplication
1412    of the 0/1 mask (clipped back to 0/1 so the counts stay small).
1413    """
1414    step = sliding_window_mask(n, w).astype(np.int64)
1415    heard = np.eye(n, dtype=np.int64)  # before any layer, each token knows only itself
1416    for _ in range(layers):
1417        heard = np.minimum(step @ heard, 1)
1418    return heard.astype(bool)

Boolean (n, n): True where token j can influence token i's output after layers windowed layers.

One layer moves information along the mask's edges; stacking layers is following edges several times, which is repeated matrix multiplication of the 0/1 mask (clipped back to 0/1 so the counts stay small).

def reach(w: int, layers: int) -> int: on GitHub
1421def reach(w: int, layers: int) -> int:
1422    """How many positions back information can travel: each layer adds w − 1."""
1423    return layers * (w - 1)

How many positions back information can travel: each layer adds w − 1.

def sliding_window_decode( Q: numpy.ndarray, K: numpy.ndarray, V: numpy.ndarray, w: int) -> tuple[numpy.ndarray, int]: on GitHub
1426def sliding_window_decode(Q: np.ndarray, K: np.ndarray, V: np.ndarray, w: int) -> tuple[np.ndarray, int]:
1427    """Generate token by token with a rolling KV buffer that holds at most w entries.
1428
1429    Returns (outputs (n, d_v), the largest the buffer ever got). The outputs
1430    equal masked attention with `sliding_window_mask`; the memory never grows past w.
1431    """
1432    keys: deque = deque(maxlen=w)  # appending the (w+1)-th entry silently drops the oldest
1433    values: deque = deque(maxlen=w)
1434    outputs, largest = [], 0
1435    for q, k, v in zip(Q, K, V):
1436        keys.append(k)
1437        values.append(v)
1438        largest = max(largest, len(keys))
1439        weights = softmax(np.array(keys) @ q / np.sqrt(len(q)))  # one row of attention, over the buffer only
1440        outputs.append(weights @ np.array(values))
1441    return np.array(outputs), largest

Generate token by token with a rolling KV buffer that holds at most w entries.

Returns (outputs (n, d_v), the largest the buffer ever got). The outputs equal masked attention with sliding_window_mask; the memory never grows past w.

def global_local_mask(n: int, w: int, global_tokens: tuple[int, ...] = (0,)) -> numpy.ndarray: on GitHub
1449def global_local_mask(n: int, w: int, global_tokens: tuple[int, ...] = (0,)) -> np.ndarray:
1450    """A sliding window plus a few global tokens that everyone reads and that read everyone (causally)."""
1451    mask = sliding_window_mask(n, w)
1452    g = list(global_tokens)
1453    mask[:, g] = True  # every token may read the global tokens...
1454    mask[g, :] = True  # ...and the global tokens may read every token...
1455    return mask & causal_mask(n)  # ...but nobody reads the future

A sliding window plus a few global tokens that everyone reads and that read everyone (causally).

def strided_mask(n: int, stride: int) -> numpy.ndarray: on GitHub
1458def strided_mask(n: int, stride: int) -> np.ndarray:
1459    """The Sparse Transformer's strided pattern: the last `stride` tokens, plus every stride-th token before them."""
1460    back = np.arange(n)[:, None] - np.arange(n)[None, :]
1461    return (back >= 0) & ((back < stride) | (back % stride == 0))

The Sparse Transformer's strided pattern: the last stride tokens, plus every stride-th token before them.

def feature_map(x: numpy.ndarray) -> numpy.ndarray: on GitHub
1469def feature_map(x: np.ndarray) -> np.ndarray:
1470    """φ(x) = elu(x) + 1: x + 1 for positive x, e^x otherwise. Always positive, so weights never go negative."""
1471    x = np.asarray(x, dtype=float)
1472    return np.where(x > 0, x + 1.0, np.exp(np.minimum(x, 0.0)))  # the minimum keeps exp from overflowing on the unused branch

φ(x) = elu(x) + 1: x + 1 for positive x, e^x otherwise. Always positive, so weights never go negative.

def linear_attention_weights(Q: numpy.ndarray, K: numpy.ndarray) -> numpy.ndarray: on GitHub
1475def linear_attention_weights(Q: np.ndarray, K: np.ndarray) -> np.ndarray:
1476    """The (n, n) causal weights linear attention implies: φ(q_i)·φ(k_j), each row divided by its total."""
1477    scores = feature_map(Q) @ feature_map(K).T  # positive "similarities", no exponential
1478    scores = np.where(causal_mask(len(Q)), scores, 0.0)  # the future contributes nothing
1479    return scores / scores.sum(axis=1, keepdims=True)

The (n, n) causal weights linear attention implies: φ(q_i)·φ(k_j), each row divided by its total.

def linear_attention_quadratic(Q: numpy.ndarray, K: numpy.ndarray, V: numpy.ndarray) -> numpy.ndarray: on GitHub
1482def linear_attention_quadratic(Q: np.ndarray, K: np.ndarray, V: np.ndarray) -> np.ndarray:
1483    """Linear attention written the slow way: build all n × n weights, then blend the values."""
1484    return linear_attention_weights(Q, K) @ V

Linear attention written the slow way: build all n × n weights, then blend the values.

def linear_attention_recurrent( Q: numpy.ndarray, K: numpy.ndarray, V: numpy.ndarray) -> tuple[numpy.ndarray, numpy.ndarray, numpy.ndarray]: on GitHub
1487def linear_attention_recurrent(Q: np.ndarray, K: np.ndarray, V: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
1488    """Linear attention the fast way: two running sums, updated once per token.
1489
1490    S (d_k, d_v) accumulates φ(k_j) v_jᵀ and z (d_k,) accumulates φ(k_j).
1491    Token i reads φ(q_i)ᵀ S / φ(q_i)ᵀ z. Returns (outputs, final S, final z):
1492    the state is the same size at token 3 and at token 3 million.
1493    """
1494    S = np.zeros((K.shape[1], V.shape[1]))
1495    z = np.zeros(K.shape[1])
1496    out = np.zeros((len(Q), V.shape[1]))
1497    for i, (q, k, v) in enumerate(zip(Q, K, V)):
1498        phi_k = feature_map(k)
1499        S = S + np.outer(phi_k, v)  # pour this token's value into the pot, flavoured by its key
1500        z = z + phi_k  # and remember how much flavour went in, to normalise later
1501        phi_q = feature_map(q)
1502        out[i] = (phi_q @ S) / (phi_q @ z)
1503    return out, S, z

Linear attention the fast way: two running sums, updated once per token.

S (d_k, d_v) accumulates φ(k_j) v_jᵀ and z (d_k,) accumulates φ(k_j). Token i reads φ(q_i)ᵀ S / φ(q_i)ᵀ z. Returns (outputs, final S, final z): the state is the same size at token 3 and at token 3 million.

def ssm_recurrent(x, A, B, C) -> numpy.ndarray: on GitHub
1516def ssm_recurrent(x, A, B, C) -> np.ndarray:
1517    """h_t = A h_{t-1} + B x_t, y_t = C · h_t, one step at a time (how an SSM runs during generation).
1518
1519    x: (T,) one input channel. A: (N, N) or a number; B, C: (N,) or numbers.
1520    """
1521    A, B, C = _ssm_params(A, B, C)
1522    h = np.zeros(len(B))  # the state: N numbers, whatever the sequence length
1523    y = []
1524    for x_t in np.asarray(x, float):
1525        h = A @ h + B * x_t
1526        y.append(float(C @ h))
1527    return np.array(y)

h_t = A h_{t-1} + B x_t, y_t = C · h_t, one step at a time (how an SSM runs during generation).

x: (T,) one input channel. A: (N, N) or a number; B, C: (N,) or numbers.

def ssm_kernel(A, B, C, length: int) -> numpy.ndarray: on GitHub
1530def ssm_kernel(A, B, C, length: int) -> np.ndarray:
1531    """The convolution kernel (C B, C A B, C A² B, ...): how much an input k steps ago still counts."""
1532    A, B, C = _ssm_params(A, B, C)
1533    kernel, A_power_B = [], B.copy()
1534    for _ in range(length):
1535        kernel.append(float(C @ A_power_B))
1536        A_power_B = A @ A_power_B
1537    return np.array(kernel)

The convolution kernel (C B, C A B, C A² B, ...): how much an input k steps ago still counts.

def ssm_convolution(x, A, B, C) -> numpy.ndarray: on GitHub
1540def ssm_convolution(x, A, B, C) -> np.ndarray:
1541    """The same outputs as `ssm_recurrent`, computed as one convolution (how a fixed SSM trains in parallel)."""
1542    x = np.asarray(x, float)
1543    kernel = ssm_kernel(A, B, C, len(x))
1544    # y_t = Σ_k kernel_k · x_{t−k}; np.convolve does exactly that sum, and we keep the first T outputs.
1545    return np.convolve(x, kernel)[: len(x)]

The same outputs as ssm_recurrent, computed as one convolution (how a fixed SSM trains in parallel).

def discretize(delta: float, a: float = -1.0, b: float = 1.0) -> tuple[float, float]: on GitHub
1548def discretize(delta: float, a: float = -1.0, b: float = 1.0) -> tuple[float, float]:
1549    """Turn a step size Δ into this step's keep and write factors (zero-order hold, one channel).
1550
1551    Ā = e^(Δa), B̄ = (e^(Δa) − 1) / a · b. With a = −1 and b = 1 they add up to 1:
1552    a big Δ overwrites the state, a tiny Δ leaves it alone.
1553    """
1554    A_bar = float(np.exp(delta * a))
1555    return A_bar, float((A_bar - 1.0) / a * b)

Turn a step size Δ into this step's keep and write factors (zero-order hold, one channel).

Ā = e^(Δa), B̄ = (e^(Δa) − 1) / a · b. With a = −1 and b = 1 they add up to 1: a big Δ overwrites the state, a tiny Δ leaves it alone.

def softplus(x): on GitHub
1558def softplus(x):
1559    """log(1 + e^x): a smooth ramp that is always positive, so a step size can never go negative."""
1560    return np.logaddexp(0.0, x)  # log(e^0 + e^x) without overflow

log(1 + e^x): a smooth ramp that is always positive, so a step size can never go negative.

def step_sizes(markers, sharpness: float = 10.0, bias: float = -5.0) -> numpy.ndarray: on GitHub
1563def step_sizes(markers, sharpness: float = 10.0, bias: float = -5.0) -> np.ndarray:
1564    """Selectivity in miniature: Δ_t = softplus(sharpness · marker_t + bias), read off each token itself.
1565
1566    Marked tokens get Δ ≈ 5 (write), unmarked ones Δ ≈ 0.007 (ignore).
1567    Mamba learns this map; here it is set by hand so the numbers are checkable.
1568    """
1569    return softplus(sharpness * np.asarray(markers, float) + bias)

Selectivity in miniature: Δ_t = softplus(sharpness · marker_t + bias), read off each token itself.

Marked tokens get Δ ≈ 5 (write), unmarked ones Δ ≈ 0.007 (ignore). Mamba learns this map; here it is set by hand so the numbers are checkable.

def selective_ssm( x, delta, a: float = -1.0, b: float = 1.0, c: float = 1.0) -> numpy.ndarray: on GitHub
1572def selective_ssm(x, delta, a: float = -1.0, b: float = 1.0, c: float = 1.0) -> np.ndarray:
1573    """One-channel selective SSM: every step discretizes with its own Δ_t, so every step has its own Ā_t and B̄_t."""
1574    h, ys = 0.0, []
1575    for x_t, d_t in zip(np.asarray(x, float), np.asarray(delta, float)):
1576        A_bar, B_bar = discretize(d_t, a, b)
1577        h = A_bar * h + B_bar * x_t
1578        ys.append(c * h)
1579    return np.array(ys)

One-channel selective SSM: every step discretizes with its own Δ_t, so every step has its own Ā_t and B̄_t.

RECALL_VALUES = array([ 0.3, -0.5, 7. , 0.8, -0.2, 0.6, -0.9, 0.4, 0.1, -0.3, 0.5, -0.6])
RECALL_MARKERS = array([0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0])
def selective_recall() -> float: on GitHub
1587def selective_recall() -> float:
1588    """The selective SSM's state after the recall task: how much of the marked 7 survived."""
1589    return float(selective_ssm(RECALL_VALUES, step_sizes(RECALL_MARKERS))[-1])

The selective SSM's state after the recall task: how much of the marked 7 survived.

def time_invariant_recall(delta: float) -> float: on GitHub
1592def time_invariant_recall(delta: float) -> float:
1593    """The same task with one fixed Δ for every token (a time-invariant SSM, like S4)."""
1594    return float(selective_ssm(RECALL_VALUES, np.full(len(RECALL_VALUES), delta))[-1])

The same task with one fixed Δ for every token (a time-invariant SSM, like S4).

def parallel_scan(a, b) -> tuple[numpy.ndarray, int]: on GitHub
1597def parallel_scan(a, b) -> tuple[np.ndarray, int]:
1598    """Every state of h_t = a_t h_{t-1} + b_t (h before the start = 0) in ceil(log2 T) rounds.
1599
1600    Each step is the map h -> a h + b. Two steps in a row are again such a map:
1601    first (a1, b1) then (a2, b2) is (a1·a2, a2·b1 + b2). In round r every
1602    position combines with the one 2^r places earlier, all positions at once,
1603    so on parallel hardware the loop over T steps becomes log2 T rounds.
1604    Returns (states (T,), rounds).
1605    """
1606    a, b = np.array(a, float), np.array(b, float)
1607    shift, rounds = 1, 0
1608    while shift < len(a):
1609        # The step `shift` places earlier; before the start it is the do-nothing map (1, 0).
1610        a_prev = np.concatenate([np.ones(shift), a[:-shift]])
1611        b_prev = np.concatenate([np.zeros(shift), b[:-shift]])
1612        a, b = a_prev * a, a * b_prev + b  # both right-hand sides use the old a
1613        shift, rounds = shift * 2, rounds + 1
1614    return b, rounds

Every state of h_t = a_t h_{t-1} + b_t (h before the start = 0) in ceil(log2 T) rounds.

Each step is the map h -> a h + b. Two steps in a row are again such a map: first (a1, b1) then (a2, b2) is (a1·a2, a2·b1 + b2). In round r every position combines with the one 2^r places earlier, all positions at once, so on parallel hardware the loop over T steps becomes log2 T rounds. Returns (states (T,), rounds).

def ssm_state_bytes(channels: int, state_size: int, bits: int = 16) -> int: on GitHub
1617def ssm_state_bytes(channels: int, state_size: int, bits: int = 16) -> int:
1618    """Memory an SSM layer carries between tokens: channels × state numbers, at any context length."""
1619    return channels * state_size * bits // 8

Memory an SSM layer carries between tokens: channels × state numbers, at any context length.

def hybrid_cache_bytes( context: int, layers: int = 32, attention_layers: int = 4, kv_heads: int = 8, head_dim: int = 128, channels: int = 4096, state_size: int = 16, bits: int = 16) -> int: on GitHub
1622def hybrid_cache_bytes(
1623    context: int, layers: int = LAYERS, attention_layers: int = 4, kv_heads: int = KV_HEADS, head_dim: int = HEAD_DIM,
1624    channels: int = 4096, state_size: int = 16, bits: int = 16,
1625) -> int:
1626    """A hybrid stack: attention layers keep a growing KV cache, SSM layers a fixed state."""
1627    attention = attention_layers * context * kv_cache_bytes_per_token(1, kv_heads, head_dim, bits)
1628    ssm = (layers - attention_layers) * ssm_state_bytes(channels, state_size, bits)
1629    return attention + ssm

A hybrid stack: attention layers keep a growing KV cache, SSM layers a fixed state.

def latent_kv_bytes_per_token(layers: int, latent_dim: int, rope_dim: int = 0, bits: int = 16) -> int: on GitHub
1637def latent_kv_bytes_per_token(layers: int, latent_dim: int, rope_dim: int = 0, bits: int = 16) -> int:
1638    """Latent KV cache per token: one latent (plus a small position-carrying key) per layer."""
1639    return layers * (latent_dim + rope_dim) * bits // 8

Latent KV cache per token: one latent (plus a small position-carrying key) per layer.

WORKED_X = array([1., 0., 2., 1.])
WORKED_W_DOWN = array([[1., 0.], [0., 1.], [1., 0.], [0., 1.]])
WORKED_W_UK = array([[1., 0., 1., 0.], [0., 1., 0., 1.]])
def latent_kv_worked_example() -> dict[str, numpy.ndarray]: on GitHub
1648def latent_kv_worked_example() -> dict[str, np.ndarray]:
1649    """x = (1, 0, 2, 1) is cached as the latent (3, 1), which expands to the key (3, 1, 3, 1) when needed."""
1650    latent = WORKED_X @ WORKED_W_DOWN
1651    return dict(latent=latent, key=latent @ WORKED_W_UK)

x = (1, 0, 2, 1) is cached as the latent (3, 1), which expands to the key (3, 1, 3, 1) when needed.

class LatentKVAttention: on GitHub
1654class LatentKVAttention:
1655    """Multi-head attention that caches one small latent per token instead of every head's keys and values.
1656
1657    The latent c = x W_down (d_latent numbers) is all that is stored. Keys
1658    and values are rebuilt from it: K = C W_uk and V = C W_uv, for all heads.
1659    Weights are random: we are studying the mechanics, not training.
1660    """
1661
1662    def __init__(self, d_model: int, n_heads: int, d_head: int, d_latent: int, seed: int = 0):
1663        rng = np.random.default_rng(seed)
1664        self.n_heads, self.d_head, self.d_latent = n_heads, d_head, d_latent
1665        s, s_latent = 1 / np.sqrt(d_model), 1 / np.sqrt(d_latent)  # keep activations near unit size
1666        self.W_q = rng.normal(0, s, (d_model, n_heads * d_head))
1667        self.W_down = rng.normal(0, s, (d_model, d_latent))  # squeeze: what gets cached
1668        self.W_uk = rng.normal(0, s_latent, (d_latent, n_heads * d_head))  # expand to every head's keys
1669        self.W_uv = rng.normal(0, s_latent, (d_latent, n_heads * d_head))  # and to every head's values
1670
1671    def _heads(self, M: np.ndarray) -> np.ndarray:
1672        # (n, H·d_head) -> (H, n, d_head): one slab per head.
1673        return M.reshape(M.shape[0], self.n_heads, self.d_head).transpose(1, 0, 2)
1674
1675    def _merge(self, M: np.ndarray) -> np.ndarray:
1676        # (H, n, d_head) -> (n, H·d_head): heads side by side again.
1677        return M.transpose(1, 0, 2).reshape(M.shape[1], self.n_heads * self.d_head)
1678
1679    def compress(self, X: np.ndarray) -> np.ndarray:
1680        """(n, d_model) -> (n, d_latent): the whole KV cache for these tokens."""
1681        return X @ self.W_down
1682
1683    def attend(self, X: np.ndarray, C: np.ndarray) -> np.ndarray:
1684        """Causal attention the direct way: expand the cached latents into keys and values, then attend."""
1685        Q, K, V = self._heads(X @ self.W_q), self._heads(C @ self.W_uk), self._heads(C @ self.W_uv)
1686        out, _ = scaled_dot_product_attention(Q, K, V, mask=causal_mask(len(X)))
1687        return self._merge(out)
1688
1689    def attend_absorbed(self, X: np.ndarray, C: np.ndarray) -> np.ndarray:
1690        """The same result without ever building K or V: fold W_uk into the query and W_uv into the output.
1691
1692        q · (c W_uk) = (q W_ukᵀ) · c, and Σ w_j (c_j W_uv) = (Σ w_j c_j) W_uv.
1693        """
1694        Q = self._heads(X @ self.W_q)  # (H, n, d_head)
1695        W_uk = self.W_uk.reshape(self.d_latent, self.n_heads, self.d_head).transpose(1, 0, 2)  # (H, d_latent, d_head)
1696        W_uv = self.W_uv.reshape(self.d_latent, self.n_heads, self.d_head).transpose(1, 0, 2)
1697        Q_latent = Q @ W_uk.transpose(0, 2, 1)  # (H, n, d_latent): each query translated into latent space
1698        scores = Q_latent @ C.T / np.sqrt(self.d_head)  # scores against the latents directly
1699        scores = np.where(causal_mask(len(X)), scores, -np.inf)
1700        mixed = softmax(scores) @ C  # (H, n, d_latent): blend the latents first...
1701        return self._merge(mixed @ W_uv)  # ...then expand once per query, not once per cached token
1702
1703    def cached_numbers_per_token(self) -> int:
1704        return self.d_latent
1705
1706    def full_numbers_per_token(self) -> int:
1707        """What ordinary multi-head attention would cache: a key and a value for every head."""
1708        return 2 * self.n_heads * self.d_head

Multi-head attention that caches one small latent per token instead of every head's keys and values.

The latent c = x W_down (d_latent numbers) is all that is stored. Keys and values are rebuilt from it: K = C W_uk and V = C W_uv, for all heads. Weights are random: we are studying the mechanics, not training.

LatentKVAttention( d_model: int, n_heads: int, d_head: int, d_latent: int, seed: int = 0) on GitHub
1662    def __init__(self, d_model: int, n_heads: int, d_head: int, d_latent: int, seed: int = 0):
1663        rng = np.random.default_rng(seed)
1664        self.n_heads, self.d_head, self.d_latent = n_heads, d_head, d_latent
1665        s, s_latent = 1 / np.sqrt(d_model), 1 / np.sqrt(d_latent)  # keep activations near unit size
1666        self.W_q = rng.normal(0, s, (d_model, n_heads * d_head))
1667        self.W_down = rng.normal(0, s, (d_model, d_latent))  # squeeze: what gets cached
1668        self.W_uk = rng.normal(0, s_latent, (d_latent, n_heads * d_head))  # expand to every head's keys
1669        self.W_uv = rng.normal(0, s_latent, (d_latent, n_heads * d_head))  # and to every head's values
W_q
W_down
W_uk
W_uv
def compress(self, X: numpy.ndarray) -> numpy.ndarray: on GitHub
1679    def compress(self, X: np.ndarray) -> np.ndarray:
1680        """(n, d_model) -> (n, d_latent): the whole KV cache for these tokens."""
1681        return X @ self.W_down

(n, d_model) -> (n, d_latent): the whole KV cache for these tokens.

def attend(self, X: numpy.ndarray, C: numpy.ndarray) -> numpy.ndarray: on GitHub
1683    def attend(self, X: np.ndarray, C: np.ndarray) -> np.ndarray:
1684        """Causal attention the direct way: expand the cached latents into keys and values, then attend."""
1685        Q, K, V = self._heads(X @ self.W_q), self._heads(C @ self.W_uk), self._heads(C @ self.W_uv)
1686        out, _ = scaled_dot_product_attention(Q, K, V, mask=causal_mask(len(X)))
1687        return self._merge(out)

Causal attention the direct way: expand the cached latents into keys and values, then attend.

def attend_absorbed(self, X: numpy.ndarray, C: numpy.ndarray) -> numpy.ndarray: on GitHub
1689    def attend_absorbed(self, X: np.ndarray, C: np.ndarray) -> np.ndarray:
1690        """The same result without ever building K or V: fold W_uk into the query and W_uv into the output.
1691
1692        q · (c W_uk) = (q W_ukᵀ) · c, and Σ w_j (c_j W_uv) = (Σ w_j c_j) W_uv.
1693        """
1694        Q = self._heads(X @ self.W_q)  # (H, n, d_head)
1695        W_uk = self.W_uk.reshape(self.d_latent, self.n_heads, self.d_head).transpose(1, 0, 2)  # (H, d_latent, d_head)
1696        W_uv = self.W_uv.reshape(self.d_latent, self.n_heads, self.d_head).transpose(1, 0, 2)
1697        Q_latent = Q @ W_uk.transpose(0, 2, 1)  # (H, n, d_latent): each query translated into latent space
1698        scores = Q_latent @ C.T / np.sqrt(self.d_head)  # scores against the latents directly
1699        scores = np.where(causal_mask(len(X)), scores, -np.inf)
1700        mixed = softmax(scores) @ C  # (H, n, d_latent): blend the latents first...
1701        return self._merge(mixed @ W_uv)  # ...then expand once per query, not once per cached token

The same result without ever building K or V: fold W_uk into the query and W_uv into the output.

q · (c W_uk) = (q W_ukᵀ) · c, and Σ w_j (c_j W_uv) = (Σ w_j c_j) W_uv.

def cached_numbers_per_token(self) -> int: on GitHub
1703    def cached_numbers_per_token(self) -> int:
1704        return self.d_latent
def full_numbers_per_token(self) -> int: on GitHub
1706    def full_numbers_per_token(self) -> int:
1707        """What ordinary multi-head attention would cache: a key and a value for every head."""
1708        return 2 * self.n_heads * self.d_head

What ordinary multi-head attention would cache: a key and a value for every head.

def quantize_kv(M: numpy.ndarray, bits: int) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
1711def quantize_kv(M: np.ndarray, bits: int) -> tuple[np.ndarray, np.ndarray]:
1712    """Quantize cached keys or values with one scale per row (per token, per head): codes and scales."""
1713    return quantize(M, bits, per_channel=True)  # `primer.ml.inference.quantize`: per_channel means one scale per row

Quantize cached keys or values with one scale per row (per token, per head): codes and scales.

def quantized_kv_bytes_per_token( layers: int, kv_heads: int, head_dim: int, bits: int, scale_bits: int = 16) -> float: on GitHub
1716def quantized_kv_bytes_per_token(layers: int, kv_heads: int, head_dim: int, bits: int, scale_bits: int = 16) -> float:
1717    """2 (key and value) × layers × KV heads × (head_dim codes of `bits` each + one scale)."""
1718    return 2 * layers * kv_heads * (head_dim * bits / 8 + scale_bits / 8)

2 (key and value) × layers × KV heads × (head_dim codes of bits each + one scale).

def quantized_cache_error(bits: int, n: int = 64, d: int = 64, seed: int = 0) -> float: on GitHub
1721def quantized_cache_error(bits: int, n: int = 64, d: int = 64, seed: int = 0) -> float:
1722    """Relative change in causal attention output when K and V are stored at `bits` bits instead of full precision."""
1723    rng = np.random.default_rng(seed)
1724    Q, K, V = (rng.standard_normal((n, d)) for _ in range(3))
1725    exact, _ = scaled_dot_product_attention(Q, K, V, mask=causal_mask(n))
1726    K_hat, V_hat = dequantize(*quantize_kv(K, bits)), dequantize(*quantize_kv(V, bits))
1727    approx, _ = scaled_dot_product_attention(Q, K_hat, V_hat, mask=causal_mask(n))
1728    return float(np.linalg.norm(approx - exact) / np.linalg.norm(exact))

Relative change in causal attention output when K and V are stored at bits bits instead of full precision.

def cache_bytes_by_method(context: int) -> dict[str, float]: on GitHub
1731def cache_bytes_by_method(context: int) -> dict[str, float]:
1732    """Memory to hold one conversation's past, for each design, on the 32-layer running example."""
1733    return {
1734        "multi-head (32 KV heads)": kv_cache_bytes(context, LAYERS, 32, HEAD_DIM),
1735        "grouped-query (8 KV heads)": kv_cache_bytes(context, LAYERS, KV_HEADS, HEAD_DIM),
1736        "grouped-query, 4-bit": context * quantized_kv_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM, bits=4),
1737        "latent (576 per layer)": context * latent_kv_bytes_per_token(LAYERS, 512, 64),
1738        "sliding window 4k": kv_cache_bytes(min(context, 4096), LAYERS, KV_HEADS, HEAD_DIM),
1739        "hybrid, 1 attention in 8": hybrid_cache_bytes(context),
1740        "state-space": LAYERS * ssm_state_bytes(4096, 16),
1741    }

Memory to hold one conversation's past, for each design, on the 32-layer running example.

def figures() -> dict: on GitHub
1749def figures() -> dict:
1750    """Plot this lesson's data. matplotlib is imported here, and only here,
1751    so the lesson itself needs nothing beyond NumPy."""
1752    import matplotlib
1753
1754    matplotlib.use("Agg")
1755    import matplotlib.pyplot as plt
1756
1757    BLUE, RED, ORANGE, GREEN, PURPLE, MUTED = "#2563eb", "#dc2626", "#d97706", "#059669", "#7c3aed", "#9ca3af"
1758    figs = {}
1759
1760    # --- 1. The two bills at 8k, 128k and 1M tokens ------------------------------
1761    contexts = (8_192, 131_072, 1_048_576)
1762    labels = ("8k", "128k", "1M")
1763    costs = [context_cost(n) for n in contexts]
1764    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4))
1765    a1.bar(labels, [c["pairs_per_head_per_layer"] for c in costs], color=RED)
1766    a1.set_yscale("log")
1767    a1.set_ylabel("pairs per head per layer")
1768    a1.set_title("Compute: pairs scored grow with n²")
1769    gb = [c["kv_cache_bytes"] / 1e9 for c in costs]
1770    a2.bar(labels, gb, color=BLUE)
1771    for i, g in enumerate(gb):
1772        a2.text(i, g + 2, f"{g:.1f} GB", ha="center")
1773    a2.axhline(80, color=MUTED, ls="--")
1774    a2.text(-0.4, 84, "one 80 GB GPU", color="#4b5563")
1775    a2.set_ylabel("KV cache, GB")
1776    a2.set_ylim(0, 160)
1777    a2.set_title("Memory: the KV cache grows with n")
1778    for a in (a1, a2):
1779        a.set_xlabel("context length")
1780    fig.tight_layout()
1781    figs["context_cost"] = fig
1782
1783    # --- 2. Sliding window: one layer vs. what three layers can reach -----------
1784    fig, (a1, a2) = plt.subplots(1, 2, figsize=(8, 3.8))
1785    a1.imshow(sliding_window_mask(16, 4), cmap="Blues", vmin=0, vmax=1.3)
1786    a1.set_title("One layer, window w = 4")
1787    a2.imshow(receptive_field(16, 4, layers=3), cmap="Blues", vmin=0, vmax=1.3)
1788    a2.set_title("Reach after 3 layers: 3·(4 − 1) = 9 back")
1789    for a in (a1, a2):
1790        a.set_xlabel("key position (being read)")
1791        a.set_ylabel("query position (reading)")
1792        a.set_xticks([0, 5, 10, 15])
1793        a.set_yticks([0, 5, 10, 15])
1794        a.grid(False)
1795    fig.tight_layout()
1796    figs["window_reach"] = fig
1797
1798    # --- 3. Sparse masks and the pairs they score --------------------------------
1799    masks = [
1800        ("full causal", causal_mask(16)),
1801        ("window w = 4", sliding_window_mask(16, 4)),
1802        ("window + global token 0", global_local_mask(16, 4, (0,))),
1803        ("strided, stride 4", strided_mask(16, 4)),
1804    ]
1805    fig, axes = plt.subplots(1, 4, figsize=(12, 3.4))
1806    for a, (name, m) in zip(axes, masks):
1807        a.imshow(m, cmap="Blues", vmin=0, vmax=1.3)
1808        a.set_title(f"{name}\n{pairs_computed(m)} pairs")
1809        a.set_xticks([0, 5, 10, 15])
1810        a.set_yticks([0, 5, 10, 15])
1811        a.set_xlabel("key")
1812        a.grid(False)
1813    axes[0].set_ylabel("query")
1814    fig.tight_layout()
1815    figs["sparse_masks"] = fig
1816
1817    # --- 4. Softmax vs. linear attention weights on the same inputs -------------
1818    rng = np.random.default_rng(11)
1819    # Queries twice the usual size: decisive, as trained queries often are. Softmax sharpens; linear barely moves.
1820    Q, K = 2 * rng.standard_normal((10, 8)), rng.standard_normal((10, 8))
1821    _, soft = scaled_dot_product_attention(Q, K, K, mask=causal_mask(10))
1822    lin = linear_attention_weights(Q, K)
1823    fig, (a1, a2) = plt.subplots(1, 2, figsize=(8.4, 3.8))
1824    for a, w, title in ((a1, soft, "softmax attention"), (a2, lin, "linear attention, φ = elu + 1")):
1825        im = a.imshow(w, cmap="Blues", vmin=0, vmax=1)
1826        a.set_title(title)
1827        a.set_xlabel("key")
1828        a.set_ylabel("query")
1829        a.grid(False)
1830    fig.colorbar(im, ax=[a1, a2], fraction=0.03, label="attention weight")
1831    figs["linear_vs_softmax"] = fig
1832
1833    # --- 5. SSM kernels: A sets how long an input keeps counting ----------------
1834    steps = np.arange(151)
1835    fig, ax = plt.subplots(figsize=(6, 3.4))
1836    for A, color in ((0.5, RED), (0.9, ORANGE), (0.99, BLUE)):
1837        ax.plot(steps, ssm_kernel(A, 1.0, 1.0, len(steps)), color=color, label=f"A = {A}")
1838    ax.set_xlabel("steps since the input arrived (k)")
1839    ax.set_ylabel("how much it still counts, C·A^k·B")
1840    ax.set_title("The kernel: how fast a fixed SSM forgets")
1841    ax.legend(frameon=False)
1842    fig.tight_layout()
1843    figs["ssm_kernel"] = fig
1844
1845    # --- 6. Selective recall ------------------------------------------------------
1846    grid = np.geomspace(1e-3, 10, 60)
1847    best = grid[int(np.argmax([time_invariant_recall(d) for d in grid]))]
1848    t = np.arange(len(RECALL_VALUES))
1849    fig, ax = plt.subplots(figsize=(7, 3.6))
1850    ax.bar(t, RECALL_VALUES, color=MUTED, label="input (marked 7 at position 2)")
1851    ax.plot(t, selective_ssm(RECALL_VALUES, step_sizes(RECALL_MARKERS)), "o-", color=BLUE, label="selective: Δ set by each token")
1852    ax.plot(t, selective_ssm(RECALL_VALUES, np.full(len(t), best)), "o-", color=RED, label=f"best fixed Δ = {best:.2f}")
1853    ax.plot(t, selective_ssm(RECALL_VALUES, np.full(len(t), 5.0)), "o-", color=ORANGE, label="large fixed Δ = 5")
1854    ax.set_xlabel("position in the sequence")
1855    ax.set_ylabel("value / state")
1856    ax.set_title("Remembering one marked token through noise")
1857    ax.legend(frameon=False, fontsize=8, loc="center right")
1858    fig.tight_layout()
1859    figs["selective_recall"] = fig
1860
1861    # --- 7. Cache per token, by design ------------------------------------------
1862    designs = [
1863        ("multi-head, 32 KV heads", kv_cache_bytes_per_token(LAYERS, 32, HEAD_DIM)),
1864        ("grouped-query, 8 KV heads", kv_cache_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM)),
1865        ("latent, 512 + 64 numbers", latent_kv_bytes_per_token(LAYERS, 512, 64)),
1866        ("grouped-query 8, 4-bit", quantized_kv_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM, bits=4)),
1867        ("multi-query, 1 KV head", kv_cache_bytes_per_token(LAYERS, 1, HEAD_DIM)),
1868    ]
1869    fig, ax = plt.subplots(figsize=(7, 3.2))
1870    kb = [b / 1024 for _, b in designs]
1871    ax.barh([d for d, _ in designs][::-1], kb[::-1], color=[GREEN, ORANGE, PURPLE, BLUE, RED])
1872    for i, v in enumerate(kb[::-1]):
1873        ax.text(v + 5, i, f"{v:.0f} KB", va="center")
1874    ax.set_xlabel("KV cache per token, all 32 layers (KB)")
1875    ax.set_xlim(0, 600)
1876    ax.set_title("Shrinking the cache per token")
1877    fig.tight_layout()
1878    figs["kv_per_token"] = fig
1879
1880    # --- 8. Memory vs. context, every design ------------------------------------
1881    ns = np.geomspace(1_000, 1_048_576, 80).astype(int)
1882    rows = [cache_bytes_by_method(int(n)) for n in ns]
1883    fig, ax = plt.subplots(figsize=(7.5, 4.2))
1884    for name in rows[0]:
1885        ax.loglog(ns, [r[name] / 1e9 for r in rows], label=name)
1886    ax.axhline(80, color=MUTED, ls="--")
1887    ax.text(1_200, 95, "80 GB GPU", color="#4b5563")
1888    ax.set_xlabel("context length n (tokens)")
1889    ax.set_ylabel("cache for one conversation (GB)")
1890    ax.set_title("What each design must keep, as context grows")
1891    ax.legend(frameon=False, fontsize=8, loc="lower right")
1892    fig.tight_layout()
1893    figs["memory_vs_context"] = fig
1894
1895    return figs

Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.

def demo() -> None: on GitHub
1903def demo() -> None:
1904    banner("1. Why long context is expensive")
1905    table(
1906        ["context", "pairs per head per layer", "KV cache (GB)"],
1907        [(f"{c['n']:,}", f"{c['pairs_per_head_per_layer']:,}", c["kv_cache_bytes"] / 1e9)
1908         for c in (context_cost(n) for n in (8_192, 131_072, 1_048_576))],
1909        floatfmt=".1f",
1910    )
1911    takeaway("Pairs grow with n², the cache with n. At 1M tokens the cache alone overflows an 80 GB GPU.")
1912
1913    banner("2. Sliding-window attention")
1914    matrix("window w = 3 over 8 tokens (1 = score computed)", sliding_window_mask(8, 3).astype(int))
1915    say(
1916        f"""
1917        {pairs_computed(sliding_window_mask(8, 3))} pairs instead of {pairs_computed(causal_mask(8))}.
1918        After three layers, token 7 can hear token 1:
1919        {bool(receptive_field(8, 3, 3)[7, 1])}; token 0: {bool(receptive_field(8, 3, 3)[7, 0])}.
1920        Mistral 7B's 32 layers of w = 4,096 reach {reach(4096, 32):,} tokens back.
1921        """
1922    )
1923    rng = np.random.default_rng(0)
1924    Q, K, V = (rng.standard_normal((50, 8)) for _ in range(3))
1925    streamed, largest = sliding_window_decode(Q, K, V, w=4)
1926    masked, _ = scaled_dot_product_attention(Q, K, V, mask=sliding_window_mask(50, 4))
1927    say(f"Decoding 50 tokens with a rolling buffer: largest buffer {largest}, same outputs: {np.allclose(streamed, masked)}.")
1928    takeaway("A window caps both bills; depth still carries information L·(w − 1) positions back.")
1929
1930    banner("3. Sparse patterns: count the pairs")
1931    table(
1932        ["pattern (16 tokens)", "pairs scored"],
1933        [
1934            ("full causal", pairs_computed(causal_mask(16))),
1935            ("window w = 4", pairs_computed(sliding_window_mask(16, 4))),
1936            ("window + global token 0", pairs_computed(global_local_mask(16, 4, (0,)))),
1937            ("strided, stride 4", pairs_computed(strided_mask(16, 4))),
1938        ],
1939    )
1940    table(
1941        ["tokens", "full causal", "window 64", "ratio"],
1942        [(n, pairs_computed(causal_mask(n)), pairs_computed(sliding_window_mask(n, 64)),
1943          pairs_computed(causal_mask(n)) / pairs_computed(sliding_window_mask(n, 64))) for n in (256, 1024, 4096)],
1944        floatfmt=".1f",
1945    )
1946    takeaway("With a fixed window the pairs grow with n, not n²; global tokens keep everyone two hops apart.")
1947
1948    banner("4. Linear attention: the pot")
1949    Qw = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 0.0]])
1950    Kw = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]])
1951    Vw = np.array([[2.0], [4.0], [6.0]])
1952    out, S, z = linear_attention_recurrent(Qw, Kw, Vw)
1953    soft, _ = scaled_dot_product_attention(Qw, Kw, Vw, mask=causal_mask(3), scale=False)
1954    say(
1955        f"""
1956        Keys (1,0), (0,1), (1,1); values 2, 4, 6; the third query is (1,0).
1957        The pot S = {S[:, 0].tolist()}, the total z = {z.tolist()}, so the output is
1958        {out[2, 0]:.3f} (= 62/15). The same inputs through softmax attention give {soft[2, 0]:.3f}.
1959        """
1960    )
1961    Q, K, V = (rng.standard_normal((200, 8)) for _ in range(3))
1962    fast, S_big, z_big = linear_attention_recurrent(Q, K, V)
1963    same = np.allclose(fast, linear_attention_quadratic(Q, K, V))
1964    say(f"On 200 random tokens the running sum equals the n × n form: {same}. The state is {S_big.shape} plus {z_big.shape}, at any length.")
1965    takeaway("Swap e^(q·k) for φ(q)·φ(k) and attention becomes a running sum: O(n), fixed state, blurrier focus.")
1966
1967    banner("5. State-space models")
1968    say(
1969        f"""
1970        A = 0.5, inputs (1, 0, 0, 2). Recurrence: {ssm_recurrent([1, 0, 0, 2], 0.5, 1.0, 1.0).tolist()}.
1971        Kernel: {ssm_kernel(0.5, 1.0, 1.0, 4).tolist()}. Convolution: {ssm_convolution([1, 0, 0, 2], 0.5, 1.0, 1.0).tolist()}.
1972        """
1973    )
1974    table(
1975        ["step size Δ", "keep Ā", "write B̄"],
1976        [(d, *discretize(d)) for d in (0.01, 0.1, 1.0, 5.0)],
1977    )
1978    grid = np.geomspace(1e-3, 10, 60)
1979    say(
1980        f"""
1981        Recall task (a marked 7, then nine noise tokens): the selective SSM ends at
1982        {selective_recall():.2f}; the best fixed step ends at
1983        {max(time_invariant_recall(d) for d in grid):.2f}.
1984        """
1985    )
1986    states, rounds = parallel_scan(np.full(1024, 0.99), rng.standard_normal(1024))
1987    say(f"A parallel scan over 1,024 steps finishes in {rounds} rounds.")
1988    table(
1989        ["context", "all-attention cache (GB)", "hybrid 4 of 32 (GB)", "pure SSM state (MB)"],
1990        [(f"{n:,}", context_cost(n)["kv_cache_bytes"] / 1e9, hybrid_cache_bytes(n) / 1e9, LAYERS * ssm_state_bytes(4096, 16) / 1e6)
1991         for n in (8_192, 131_072, 1_048_576)],
1992        floatfmt=".2f",
1993    )
1994    takeaway("A selective SSM decides per token what to keep; its memory never grows, and a few attention layers restore exact recall.")
1995
1996    banner("6. Compressing the KV cache")
1997    worked = latent_kv_worked_example()
1998    say(f"Latent example: x = (1, 0, 2, 1) is cached as {worked['latent'].tolist()} and expands to the key {worked['key'].tolist()}.")
1999    layer = LatentKVAttention(d_model=64, n_heads=8, d_head=8, d_latent=16)
2000    X = rng.standard_normal((12, 64))
2001    C = layer.compress(X)
2002    say(
2003        f"""
2004        A latent layer with 8 heads of width 8 caches {layer.cached_numbers_per_token()} numbers per token
2005        instead of {layer.full_numbers_per_token()}. Expanded and absorbed attention agree:
2006        {np.allclose(layer.attend(X, C), layer.attend_absorbed(X, C))}.
2007        """
2008    )
2009    table(
2010        ["cache precision", "output error", "bytes per token (8B example)"],
2011        [(f"{b}-bit", quantized_cache_error(b), f"{quantized_kv_bytes_per_token(LAYERS, KV_HEADS, HEAD_DIM, b, scale_bits=0 if b == 16 else 16):,.0f}")
2012         for b in (16, 8, 4)],
2013    )
2014    table(
2015        ["design", "cache at 1M tokens (GB)"],
2016        [(name, b / 1e9) for name, b in cache_bytes_by_method(1_048_576).items()],
2017        floatfmt=".3f",
2018    )
2019    takeaway("Share heads, squeeze into a latent, cut bits, or bound the window: each moves the memory line down or flattens it.")