At a glance
Key takeaways
- 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.
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:
- Look at fewer tokens: sliding-window and sparse attention.
- Summarize the past into a fixed-size state: linear attention and state-space models (Mamba), which run like a recurrent network.
- Store the past more compactly: fewer key/value heads, a small latent vector per token, fewer bits per number.
Figure 1 · Diagram
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
Chapter 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.
Figure 2 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| context length: how many tokens are in play | 131,072 | |
| times : every query paired with every key (causal masking halves it, which does not change how it grows) | 17,179,869,184 | |
| one key and one value per token | ||
| layers, each with its own cache | 32 | |
| key/value heads per layer | 8 | |
| numbers per head | 128 | |
| 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
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]
Figure 3 · Drawn from the lesson's code
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
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.
Chapter 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.
Figure 4 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the mask: 1 if query may read key , 0 if not | row 5: keys 3, 4, 5 | |
| , | the query's position and the key's position | |
| how far back the key is; negative means the future | 0, 1, 2 allowed | |
| the window: how many tokens each query reads, itself included | 3 | |
| "cases": use the line whose condition holds | ||
| 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
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
Figure 5 · Drawn from the lesson's code
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
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.
Figure 6 · Diagram
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"]
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.
Chapter 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.
Figure 7 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| 1 if query may read key | ||
| how far back key is | for : 0 to 13 | |
| the stride: the local width, and the jump between long-range keys | 4 | |
| "modulo": the remainder after dividing; 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
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
Figure 8 · Drawn from the lesson's code
Four 16 by 16 masks: full causal attention scores 136 pairs, the window 58, window plus a global first token 70, strided 82
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.
Chapter 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 cannot
be split into "a part that depends on q" times "a part that depends on k".
Replace it with , a dot product of transformed
vectors, and it can. Then the order of the matrix multiplies can be swapped:
. 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 (always above zero).
- φ of the keys: (2, 1), (1, 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).
- 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.
Figure 9 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the output for token | ||
| , , | query of token ; key and value of token | |
| the feature map, applied to each number: if , else ; always positive so no weight is negative | φ(1, 0) = (2, 1) | |
| add up over every token up to and including (causal) | = 1, 2, 3 | |
| the unnormalized weight, a dot product in place of | 5, 4, 6 | |
| the value laid on its side as a row | ||
| an outer product: a column times a row, giving a small table whose entry (m, c) is | (2, 1) × 2 = (4, 2) | |
| the pot after token : the sum of those tables | (20, 22) | |
| the running sum of , for the denominator | (5, 5) | |
| the query's ladle: its dot product with each column of | 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
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
Figure 10 · Drawn from the lesson's code
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
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 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.
Chapter 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.
Figure 11 · Diagram
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
primer.ml.cnn_rnn. Classic RNNs only had
the top road, which is why they trained slowly.Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the input at step (one channel) | (1, 0, 0, 2) | |
| the state after step : numbers (here = 1); | 1, 0.5, 0.25, 2.125 | |
| matrix: how the old state carries over (its size sets how fast things fade) | 0.5 | |
| how the input is written into the state | 1 | |
| how the state is read out | 1 | |
| the output at step | ||
| "the same thing, written another way" | ||
| multiplied by itself times; is "do nothing" | 0.5³ = 0.125 | |
| the kernel: how much an input steps ago still counts now | 1, 0.5, 0.25, 0.125 | |
| add over every look-back distance from 0 to |
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
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
Figure 12 · Drawn from the lesson's code
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
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 , with negative, so a big Δ makes it nearly 0 (forget) and a tiny Δ nearly 1 (keep);
- write factor , 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.
Figure 13 · Diagram
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"]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| what the model can see about token (here, just its marker, 0 or 1) | 1 for the 7 | |
| , | learned weight and offset that turn into a step size | 10, −5 |
| softplus | : a smooth ramp that is always positive, about for large and about 0 for very negative | softplus(5) = 5.007 |
| the step size for token : how much this token matters | 5.007 or 0.0067 | |
| a negative learned rate; more negative means faster forgetting | −1 | |
| how strongly inputs are written | 1 | |
| this step's keep factor ("A-bar") | 0.0067 or 0.9933 | |
| this step's write factor ("B-bar") | 0.9933 or 0.0067 | |
| , | the value written, and the state after step | 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
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
Figure 14 · Drawn from the lesson's code
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
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.
Figure 15 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| one step: "multiply the state by , then add " | (0.5, 1) | |
| "do the first step, then the second": merging two steps into one | ||
| the combined multiplier | 0.25 | |
| 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
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.
Figure 16 · Diagram
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"]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| total layers | 32 | |
| how many of them are attention layers | 4 | |
| context length | 131,072 | |
| one attention layer's cache per token (from section 1) | 4,096 bytes | |
| channels in an SSM layer | 4,096 | |
| state numbers per channel | 16 | |
| 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
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.
Chapter 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:
- Squeeze: c = x W_down = (1 + 2, 0 + 1) = (3, 1). Only these 2 numbers are cached: 4 times smaller.
- When needed, expand: k = c W_uk = (3, 1, 3, 1), both heads' keys.
- 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.
Figure 17 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | Shape / example |
|---|---|---|
| token 's vector coming into the layer | ; (1, 0, 2, 1) | |
| the learned "down" projection that squeezes | ||
| the latent: the only thing cached | ; (3, 1) | |
| , | learned "up" projections to every head's keys and values | |
| , | token 's keys and values for all heads, side by side | ; = (3, 1, 3, 1) |
| a query (one head's slice, or all heads side by side) | (1, 2, 0, 1) | |
| transposed (rows become columns) | ||
| the query translated into latent space | ; (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
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.
Figure 18 · Drawn from the lesson's code
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
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.
Figure 19 · Diagram
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"]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| how many head vectors one token stores: a key and a value per KV head per layer | 2 × 32 × 8 = 512 | |
| 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
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.
Opening
Putting it together
Figure 20 · Drawn from the lesson's code
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
| 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.
Test yourself
8 questions
Answer each one out loud or on paper before you open it. If you can explain it, you know it.
Question 1Why does doubling the context quadruple attention's compute but only double its cache?Think it through, then reveal
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.
Question 2A model uses a 4,096-token sliding window in all 32 layers. Can token 100,000 be influenced by token 1?Think it through, then reveal
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.
Question 3What does a global token do in a sparse pattern, and why is it cheap?Think it through, then reveal
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².
Question 4Why can linear attention run as a recurrence but softmax attention cannot?Think it through, then reveal
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.
Question 5What does "selective" mean in Mamba, and what problem does it fix?Think it through, then reveal
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.
Question 6If a selective SSM's parameters change every step, how does it train in parallel?Think it through, then reveal
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.
Question 7How does latent KV caching save memory without changing the attention outputs?Think it through, then reveal
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.
Question 8Why do hybrid models keep a few attention layers instead of going all-SSM?Think it through, then reveal
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.
Primary sources
The papers behind this lesson
Introduced strided and fixed sparse attention patterns, cutting attention's cost to about n√n.
The paper ↗Combined a sliding window with task-chosen global tokens to read long documents in linear time.
The paper ↗Mixed window, global and random links, and proved such sparse attention keeps the expressive power of full attention.
The paper ↗Replaced softmax with a kernel feature map, turning attention into a running sum with O(n) cost.
The paper ↗Made linear state-space layers trainable on very long sequences through the convolution view.
The paper ↗Made the state-space parameters depend on the input and trained them with a hardware-aware parallel scan.
Read the annotated companion →The paper ↗Showed that selective SSMs and a form of linear attention are two views of one computation.
The paper ↗Used sliding-window attention with a rolling-buffer cache in a strong open model.
Read the annotated companion →The paper ↗Found that models lean on the first few tokens, and that keeping them plus a window allows streaming without limit.
The paper ↗Interleaved one attention layer per seven Mamba layers to cut the KV cache while keeping recall.
The paper ↗Introduced multi-query attention, one shared key/value head for all query heads.
Read the annotated companion →The paper ↗Introduced multi-head latent attention, caching a small latent per token instead of full keys and values.
The paper ↗Quantized keys per channel and values per token, bringing the cache down to 2 bits.
The paper ↗Researcher's shelf
Further reading
- Child et al., Sparse Transformers (2019): https://arxiv.org/abs/1904.10509
- Beltagy et al., Longformer (2020): https://arxiv.org/abs/2004.05150
- Katharopoulos et al., Transformers are RNNs (2020): https://arxiv.org/abs/2006.16236
- Gu, Goel & Ré, S4 (2021): https://arxiv.org/abs/2111.00396
- Sasha Rush et al., The Annotated S4 (the S4 paper, line by line in code): https://srush.github.io/annotated-s4/
- Gu & Dao, Mamba (2023): https://arxiv.org/abs/2312.00752
- Dao & Gu, Mamba-2 (2024): https://arxiv.org/abs/2405.21060
- Lieber et al., Jamba (2024): https://arxiv.org/abs/2403.19887
- DeepSeek-AI, DeepSeek-V2 (2024): https://arxiv.org/abs/2405.04434
- Ainslie et al., GQA (2023): https://arxiv.org/abs/2305.13245
About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit c8d5c21, so the two always agree: the explanation, the code that builds it and the tests that prove it.