Attention, Part 3: Sliding-Window and Sparse Attention

Why looking at every earlier token gets too expensive, and the three ways today's models avoid it: sliding windows (Mistral), mixing local and global layers (Gemma 2, Gemma 3, gpt-oss), and letting the model pick its own tokens (DeepSeek Sparse Attention). Explained from zero, with equations, quotes from the papers, and real experiments: a lightning indexer trained on Qwen2.5-0.5B keeps quality close to full attention while each token reads only 64 of up to 2,048 earlier tokens.

Part 2 made each token store less: fewer key/value heads (MQA, GQA) or a compressed latent (MLA). But every token still looked at every earlier token.

This part attacks that second problem: how many tokens each token looks at. It is the biggest cost of long conversations, long documents and long chains of reasoning, and it is why every model that handles 128,000 tokens or more uses at least one of the ideas below.

This is a long part, so here is the map.

As in the earlier parts, every technical word gets a yellow box the first time it appears, every number comes from code I ran, and every quote links to its paper.

1. The problem: everyone looks at everyone

In full attention, token number ii compares its query with the keys of tokens 1,2,…,i1, 2, \dots, i. So for a text of TT tokens, the number of query-key comparisons in one head of one layer is:

1+2+3+⋯+T=T (T+1)21 + 2 + 3 + \dots + T = \frac{T\,(T+1)}{2}

Here is what that means in numbers:

Text length TTComparisons, full attentionComparisons, window of 1,024Full ÷ window
1,024 tokens524,800524,8001.0×
8,192 tokens33,558,5287,864,8324.3×
131,072 tokens (128K)8,590,000,128133,693,95264.3×

At 128K tokens, full attention makes 8.6 billion comparisons, in every head of every layer. A window of 1,024 tokens makes 64 times fewer.

In arithmetic operations, each comparison is a dot product of length dd (about 2d2d multiply-and-add steps), and the weighted mix of values costs about the same again. So per head and layer:

workfull≈2⋅2d⋅T(T+1)2  ≈  2 d T2workwindow≈2⋅2d⋅T W  =  4 d T W\text{work}_{\text{full}} \approx 2 \cdot 2d \cdot \frac{T(T+1)}{2} \;\approx\; 2\,d\,T^2 \qquad\qquad \text{work}_{\text{window}} \approx 2 \cdot 2d \cdot T\,W \;=\; 4\,d\,T\,W

where TT is the text length, WW the window and dd the head size. The first grows with T2T^2, the second only with TT: doubling a document doubles the windowed cost but quadruples the full cost.

There is also a memory cost while the model writes. The KV cache grows by one token at every step, and every step reads all of it. Twice the conversation means twice the memory and twice the reading per step.

The rest of this part is about skipping the comparisons that do not matter, without skipping the ones that do.

full (causal)sliding window, W = 4sparse: top-k pickedEach row is one token; each bright square is a token it may look at. Faint squares are skipped. Dark empty squares are the future.
Three patterns on 12 tokens. Full attention: every token sees all earlier tokens. Sliding window: each token sees only itself and the 3 before it. Sparse top-k: each token sees a small, chosen set (here an illustrative random pick).

2. Sliding-window attention (SWA)

The simplest fix: each token looks only at the last W tokens, itself included. Nothing older.

In Part 1 we wrote attention with a mask MM that blocks the future. Sliding-window attention is the same formula with a stricter mask:

Attention(Q,K,V)=softmax⁡ ⁣(QK⊤d+M)V,Mij={0if 0≤i−j<W−∞otherwise\text{Attention}(Q, K, V) = \operatorname{softmax}\!\left(\frac{QK^\top}{\sqrt{d}} + M\right) V, \qquad M_{ij} = \begin{cases} 0 & \text{if } 0 \le i - j < W \\ -\infty & \text{otherwise} \end{cases}

where:

  • ii is the position of the query token and jj the position of a key token;
  • i−ji - j is how far back token jj is from token ii;
  • i−j<0i - j < 0 would be the future (blocked, as before);
  • i−j≥Wi - j \ge W is too far back (newly blocked);
  • −∞-\infty becomes a weight of exactly 0 after softmax.

That is the whole idea. In code it is one extra condition on the mask:

python
def window_mask(T, W, device=None):
    """True where query i may look at key j: j <= i (causal) and i - j < W (inside the window)."""
    i = torch.arange(T, device=device)[:, None]
    j = torch.arange(T, device=device)[None, :]
    return (j <= i) & (i - j < W)


def attend(q, k, v, allowed):
    """Plain attention with a boolean 'allowed' mask. q, k, v: (..., T, d)."""
    scores = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1])
    scores = scores.masked_fill(~allowed, float('-inf'))
    return torch.softmax(scores, -1) @ v

For 10 tokens and W=4W = 4, the mask looks like this (1 = may look, 0 = may not). Each row is one token; it sees itself and the three before it:

plain text
1 0 0 0 0 0 0 0 0 0
1 1 0 0 0 0 0 0 0 0
1 1 1 0 0 0 0 0 0 0
1 1 1 1 0 0 0 0 0 0
0 1 1 1 1 0 0 0 0 0
0 0 1 1 1 1 0 0 0 0
0 0 0 1 1 1 1 0 0 0
0 0 0 0 1 1 1 1 0 0
0 0 0 0 0 1 1 1 1 0
0 0 0 0 0 0 1 1 1 1

Proof that it works

I checked three things in part3_sparse.py (24 tokens, W=6W = 6, 64-bit numbers):

plain text
1. SWA vs PyTorch SDPA with the same mask:  max |diff| = 4.4e-16
   window wider than the text vs causal:   max |diff| = 0.0e+00
   change all 18 tokens outside the window of the last token:
      last output with sliding window moves by 0.0e+00
      last output with full attention moves by 0.65
  1. My function agrees with PyTorch's built-in attention to 4.4×10−164.4 \times 10^{-16} (rounding noise).
  2. A window wider than the text is exactly ordinary causal attention (difference 0).
  3. The real test of "only the last W tokens matter": I replaced all 18 tokens outside the last token's window with random new ones. With the sliding window, the last token's output did not move at all (difference exactly 0). With full attention, it moved by 0.65.

3. "But then the model forgets everything older than W?"

Not quite, and this is the clever part. A model has many layers stacked on top of each other. Layer 2 reads the outputs of layer 1, and each of those outputs already mixed information from its own window.

input tokensafter layer 1after layer 2after layer 3Window W = 3. Each layer reaches W − 1 = 2 more tokens back: 1 → 3 → 5 → 7 tokens.So after L layers, a token can be influenced by L × (W − 1) + 1 tokens.
How the reach grows. Window W = 3. The last token after layer 3 reads 3 tokens of layer 2, each of which read 3 tokens of layer 1, and so on. By the input, 7 tokens can affect it.

Each layer lets information travel W−1W - 1 more tokens back. After LL layers:

reach=L×(W−1)+1 tokens\text{reach} = L \times (W - 1) + 1 \ \text{tokens}

The Mistral 7B paper says the same thing (counting the window with its W+1W+1 convention):

Longformer (2020) wrote the same rule in its own notation, with ℓ\ell layers and window ww:

Proof. I stacked 1 to 6 sliding-window layers (W=4W = 4, with the usual "add the input back" connection that real models use), and asked PyTorch which input tokens the last output depends on. It can tell exactly, by computing the gradient.

plain text
2. Receptive field of the last token, window W = 4
   1 layer(s): depends on  4 tokens (formula L*(W-1)+1 =  4), reaches  3 tokens back
   2 layer(s): depends on  7 tokens (formula L*(W-1)+1 =  7), reaches  6 tokens back
   3 layer(s): depends on 10 tokens (formula L*(W-1)+1 = 10), reaches  9 tokens back
   4 layer(s): depends on 13 tokens (formula L*(W-1)+1 = 13), reaches 12 tokens back
   5 layer(s): depends on 16 tokens (formula L*(W-1)+1 = 16), reaches 15 tokens back
   6 layer(s): depends on 19 tokens (formula L*(W-1)+1 = 19), reaches 18 tokens back

The measured reach matches the formula exactly, at every depth.

4. The rolling buffer cache: memory that never grows

With a window, a token never needs keys and values older than WW steps. So the KV cache can be a fixed-size ring.

Writing token 9 with a cache of W = 4 slotstok 0tok 1tok 2tok 3tok 4tok 5tok 6tok 7tok 8tok 9tokens 0 to 5: overwritten (crossed out); tokens 6 to 9: still in the cacheslot 0holds tok 8slot 1holds tok 9slot 2holds tok 6slot 3holds tok 7Rule: token i goes into slot i mod 4. Token 9 → slot 1, overwriting token 5.The cache never grows past W slots, no matter how long the text gets.
A rolling buffer with W = 4 slots while writing token 9. Token i goes into slot i mod 4. Tokens 0 to 5 have been overwritten; the cache always holds just the last 4 tokens.

Here is the decoding loop from part3_sparse.py. It writes one token at a time into a cache of only WW slots:

python
k_cache = torch.zeros(H, W, d, dtype=torch.float64)                     # the whole cache: W slots, never more
v_cache = torch.zeros(H, W, d, dtype=torch.float64)
step_out = []
for i in range(T):
    xi = x[i:i + 1]
    qi = rope((xi @ Wq).view(1, H, d).transpose(0, 1), pos[i:i + 1])
    ki = rope((xi @ Wk).view(1, H, d).transpose(0, 1), pos[i:i + 1])
    vi = (xi @ Wv).view(1, H, d).transpose(0, 1)
    slot = i % W                                                        # position i goes to slot i mod W
    k_cache[:, slot], v_cache[:, slot] = ki[:, 0], vi[:, 0]
    n = min(i + 1, W)                                                   # how many slots are filled
    step_out.append(attend(qi, k_cache[:, :n], v_cache[:, :n], torch.ones(1, n, dtype=torch.bool)))

Notice that the slots are out of order: in the picture, slot 0 holds token 8 and slot 2 holds token 6. Does that break anything? No, for two reasons:

  1. Softmax and the weighted sum do not care about order. Adding up "weight × value" over a set of tokens gives the same answer in any order.
  2. Position is already baked into each key. RoPE (explained in Part 2) rotates each key by its position before it is stored, so the key remembers where it came from, whatever slot it sits in.

Proof. I compared this token-by-token ring cache with computing all 50 tokens at once using the sliding-window mask (2 heads, RoPE on):

plain text
3. Rolling buffer cache (W = 8 slots) vs computing all 50 tokens at once: max |diff| = 1.1e-15

Identical, while the cache held 8 slots instead of 50.

How much memory does it save?

The KV cache formula from Part 2, with the number of stored tokens capped at WW:

KV cache=2×L×Hkv×dh×min⁡(n,W)×b\text{KV cache} = 2 \times L \times H_{kv} \times d_h \times \min(n, W) \times b

where nn is the number of tokens so far, and the other symbols are as in Part 2 (layers, key/value heads, head size, bytes per number).

For Mistral 7B (32 layers, 8 key/value heads, head size 128, window 4,096) at 32,768 tokens:

plain text
   Mistral 7B at 32K: full 4.00 GiB, window 4096 0.50 GiB (8x smaller)

And speed?

Every writing step reads the whole cache. With full attention that read keeps growing; with a window it stays the same size. I timed one attention step (16 heads, head size 128, 16-bit numbers) on an Apple M5 Pro GPU:

0.001.002.003.004.001,0244,09616,38465,536131,072full attention, 1,024: 0.06full attention, 4,096: 0.14full attention, 16,384: 0.52full attention, 65,536: 2.02full attention, 131,072: 3.89window 1,024, 1,024: 0.04window 1,024, 4,096: 0.03window 1,024, 16,384: 0.03window 1,024, 65,536: 0.03window 1,024, 131,072: 0.03full attention: 3.89window 1,024: 0.03tokens so far (log scale)milliseconds per step
Time for one decode step of attention. Full attention gets slower as the conversation grows, because it reads the whole cache. With a 1,024-token window, the step time stays flat.
Tokens so farFull attentionWindow of 1,024
1,0240.057 ms0.041 ms
4,0960.137 ms0.033 ms
16,3840.521 ms0.030 ms
65,5362.019 ms0.034 ms
131,0723.893 ms0.031 ms

At 128K tokens, the windowed step is about 125 times faster. (At 1,024 tokens both read the same amount; the small gap there is timing noise.)

5. Mixing local and global layers

A pure sliding-window model has the weakness from the warning above: it can never look up an old token directly. The fix that most recent models use is to mix two kinds of layers:

This idea is older than chatbots. Longformer (2020) already combined the two inside one layer:

Today's models do it layer by layer:

  • Gemma 2 (2024) alternates 1 local : 1 global, with a local window of 4,096 tokens.
  • Gemma 3 (2025) goes further: 5 local : 1 global, and shrinks the window to 1,024.
  • gpt-oss (OpenAI, 2025) alternates 1 : 1 with a tiny window of just 128 tokens.

Do the local layers hurt quality? Gemma 3 tested it:

Gemma 3 27B: 62 layers, 5 local (sliding window 1,024) then 1 global, repeatedlayer 1: local (last 1,024 tokens)layer 2: local (last 1,024 tokens)layer 3: local (last 1,024 tokens)layer 4: local (last 1,024 tokens)layer 5: local (last 1,024 tokens)layer 6: global (sees every token)layer 7: local (last 1,024 tokens)layer 8: local (last 1,024 tokens)layer 9: local (last 1,024 tokens)layer 10: local (last 1,024 tokens)layer 11: local (last 1,024 tokens)layer 12: global (sees every token)layer 13: local (last 1,024 tokens)layer 14: local (last 1,024 tokens)layer 15: local (last 1,024 tokens)layer 16: local (last 1,024 tokens)layer 17: local (last 1,024 tokens)layer 18: global (sees every token)layer 19: local (last 1,024 tokens)layer 20: local (last 1,024 tokens)layer 21: local (last 1,024 tokens)layer 22: local (last 1,024 tokens)layer 23: local (last 1,024 tokens)layer 24: global (sees every token)layer 25: local (last 1,024 tokens)layer 26: local (last 1,024 tokens)layer 27: local (last 1,024 tokens)layer 28: local (last 1,024 tokens)layer 29: local (last 1,024 tokens)layer 30: global (sees every token)layer 31: local (last 1,024 tokens)layer 32: local (last 1,024 tokens)layer 33: local (last 1,024 tokens)layer 34: local (last 1,024 tokens)layer 35: local (last 1,024 tokens)layer 36: global (sees every token)layer 37: local (last 1,024 tokens)layer 38: local (last 1,024 tokens)layer 39: local (last 1,024 tokens)layer 40: local (last 1,024 tokens)layer 41: local (last 1,024 tokens)layer 42: global (sees every token)layer 43: local (last 1,024 tokens)layer 44: local (last 1,024 tokens)layer 45: local (last 1,024 tokens)layer 46: local (last 1,024 tokens)layer 47: local (last 1,024 tokens)layer 48: global (sees every token)layer 49: local (last 1,024 tokens)layer 50: local (last 1,024 tokens)layer 51: local (last 1,024 tokens)layer 52: local (last 1,024 tokens)layer 53: local (last 1,024 tokens)layer 54: global (sees every token)layer 55: local (last 1,024 tokens)layer 56: local (last 1,024 tokens)layer 57: local (last 1,024 tokens)layer 58: local (last 1,024 tokens)layer 59: local (last 1,024 tokens)layer 60: global (sees every token)layer 61: local (last 1,024 tokens)layer 62: local (last 1,024 tokens)52 local layers: keep only the last 1,024 tokens10 global layers: keep every tokenlayer 1layer 62
The real layer pattern of Gemma 3 27B: five local layers (blue, window 1,024), then one global layer (orange), repeated. 52 local and 10 global layers in total. Hover a bar to see the layer.

Here is the published configuration (from Hugging Face), which is where those numbers come from:

json
{
  "num_hidden_layers": 62,
  "num_attention_heads": 32,
  "num_key_value_heads": 16,
  "head_dim": 128,
  "sliding_window": 1024,
  "sliding_window_pattern": 6,
  "rope_theta": 1000000.0,
  "rope_local_base_freq": 10000.0
}

sliding_window_pattern: 6 means every 6th layer is global: layers 6, 12, 18, ..., 60. That gives 10 global and 52 local layers.

The memory math

Global layers store every token; local layers store at most WW:

KV cache=2 Hkv dh b (Lglobal⋅n  +  Llocal⋅min⁡(n,W))\text{KV cache} = 2 \, H_{kv} \, d_h \, b \, \Big( L_{\text{global}} \cdot n \;+\; L_{\text{local}} \cdot \min(n, W) \Big)

where LglobalL_{\text{global}} and LlocalL_{\text{local}} are the numbers of global and local layers. For Gemma 3 27B (Hkv=16H_{kv} = 16, dh=128d_h = 128, 2 bytes per number):

plain text
   Gemma 3 27B (52 local + 10 global layers):
         1024 tokens: all global   0.48 GiB   5 local : 1 global  0.48 GiB   (1.0x smaller)
         2048 tokens: all global   0.97 GiB   5 local : 1 global  0.56 GiB   (1.7x smaller)
         4096 tokens: all global   1.94 GiB   5 local : 1 global  0.72 GiB   (2.7x smaller)
         8192 tokens: all global   3.88 GiB   5 local : 1 global  1.03 GiB   (3.8x smaller)
        16384 tokens: all global   7.75 GiB   5 local : 1 global  1.66 GiB   (4.7x smaller)
        32768 tokens: all global  15.50 GiB   5 local : 1 global  2.91 GiB   (5.3x smaller)
        65536 tokens: all global  31.00 GiB   5 local : 1 global  5.41 GiB   (5.7x smaller)
       131072 tokens: all global  62.00 GiB   5 local : 1 global 10.41 GiB   (6.0x smaller)
0.0163248641,0244,09616,38465,536131,072all layers global, 1,024: 0.5all layers global, 2,048: 1.0all layers global, 4,096: 1.9all layers global, 8,192: 3.9all layers global, 16,384: 7.8all layers global, 32,768: 16all layers global, 65,536: 31all layers global, 131,072: 625 local : 1 global, 1,024: 0.55 local : 1 global, 2,048: 0.65 local : 1 global, 4,096: 0.75 local : 1 global, 8,192: 1.05 local : 1 global, 16,384: 1.75 local : 1 global, 32,768: 2.95 local : 1 global, 65,536: 5.45 local : 1 global, 131,072: 10all layers global: 625 local : 1 global: 10tokens in the conversation (log scale)KV cache, GiB
KV cache of Gemma 3 27B for one conversation. If every layer were global, 128K tokens would need 62 GiB. With the real 5:1 pattern it needs 10.4 GiB.

Gemma 3's own measurement, for a 2B model, shows the same shape:

At the full 128K context: 62 GiB → 10.4 GiB. The savings can never pass 6.2× (62 layers ÷ 10 global layers), because the global layers still grow with the text. They become the main cost.

For gpt-oss-20b (24 layers alternating, 8 key/value heads, head size 64, window 128), at 128K tokens the cache drops from 6.00 GiB to 3.00 GiB: half the layers keep almost nothing.

The RoPE angle for pair ii of a head of size dd is

θi=base−2i/d,i=0,1,…,d2−1\theta_i = \text{base}^{-2i/d}, \qquad i = 0, 1, \dots, \tfrac{d}{2} - 1

where "base" is 10,000 normally and 1,000,000 in Gemma 3's global layers. The slowest pair (i=d/2−1i = d/2 - 1) has a period, in tokens, of about 2π⋅base2\pi \cdot \text{base}: roughly 63,000 tokens for base 10k, roughly 6.3 million for base 1M. That is why the global layers need the larger base to tell positions apart across 128K tokens.

ModelPatternWindowWhy it matters
Mistral 7Bevery layer local4,096made SWA and the rolling cache popular in open models
Gemma 21 local : 1 global4,096half the layers stay small
Gemma 35 local : 1 global1,0246× smaller cache at 128K
gpt-oss1 local : 1 global128tiny windows, half the layers almost free

6. Attention sinks strike again: a real experiment

All the models above were trained with their windows, so they learned to live with them. What happens if you take a normal model, trained with full attention, and simply force a window on it? This is what "StreamingLLM" studied, and the answer surprised people.

Remember the attention sink from Part 1: in Qwen2.5-0.5B, 68% of heads put most of their attention on the very first token, because softmax forces the weights to add up to 1 and the heads need somewhere harmless to "park" them. A sliding window cuts that first token off.

I tested this myself.

Setup (part3_qwen.py):

  • Model: Qwen2.5-0.5B, trained with full attention.
  • Text: four chunks of 2,048 tokens from Pride and Prejudice (public domain, from Project Gutenberg).
  • I score only tokens 1,024 to 2,047, where every window really cuts something off.
  • Each rule gets the same budget: each token may look at exactly 64, 256 or 1,024 tokens.
  • "4 sinks + window" means: the first 4 tokens of the text, plus the most recent (budget − 4) tokens.
  • First, a sanity check: passing my own full-attention mask gives exactly the model's normal output (max difference 0.0).
plain text
full attention: perplexity 17.39
budget   64 tokens: window only   132.70   4 sinks + window 60:  22.45
budget  256 tokens: window only    67.90   4 sinks + window 252:  18.83
budget 1024 tokens: window only   508.88   4 sinks + window 1020:  17.80

What this shows:

  • A plain window destroys the model. Perplexity jumps from 17.4 to between 68 and 509.
  • Keeping just 4 sink tokens fixes almost all of it. With 1,024 tokens of budget, 4 sinks plus a window give 17.80, very close to full attention's 17.39.
  • A bigger window does not save you if the sink is gone. The 1,024-token window without sinks was the worst of all (508.9), even though it sees 16 times more text than the 64-token window. I did not expect this. It shows the problem is not missing information: the model is thrown off by losing its "parking spot".

With SS sink tokens and a window of WW recent tokens, the cache holds a fixed

cache size=(S+W)×bytes per token\text{cache size} = (S + W) \times \text{bytes per token}

whatever the length of the conversation: 4 + 1,020 = 1,024 tokens in my 1,024-budget test.

7. Sparse attention: let the model choose

Windows are fixed patterns: they always keep the most recent tokens. But sometimes the important token is far away. A character's name introduced in chapter 1, a function defined at the top of a file, the original question at the start of a long chat.

Early sparse transformers (Child et al., 2019) used fixed patterns, like "the last few tokens plus every 64th token". The newer idea is to let the model pick the tokens that matter for each query. The problem: to know which tokens score highest, you seem to need the scores, which is the very thing you were trying to avoid computing.

DeepSeek's answer: compute the scores with a much smaller, much cheaper model first, then do the expensive attention only on the winners.

8. DeepSeek Sparse Attention (DSA)

DSA arrived with DeepSeek-V3.2 (2025). It has two pieces.

token tquerylightning indexerfew heads, small, FP8scores every earlier tokentop-kkeep the best kmain attentiononly over thek chosen tokensoutputcheap, but looks at all L tokensexpensive, but looks at only kDeepSeek-V3.2: k = 2,048. The main attention cost drops from L² to L × k; the indexer is still L², but tiny.
DeepSeek Sparse Attention. The lightning indexer is small and cheap, but it scores every earlier token. Top-k keeps the best k. The big main attention then runs only over those k tokens.

Piece 1: the lightning indexer

The index score is:

It,s=∑j=1HIwt,jI⋅ReLU⁡ ⁣(qt,jI⋅ksI)I_{t,s} = \sum_{j=1}^{H^I} w^I_{t,j} \cdot \operatorname{ReLU}\!\left(q^I_{t,j} \cdot k^I_s\right)

where:

  • tt is the current (query) token and ss an earlier token;
  • HIH^I is the number of indexer heads (a small number);
  • qt,jIq^I_{t,j} is the indexer query of token tt for indexer head jj;
  • ksIk^I_s is the indexer key of token ss (one key per token, shared by all indexer heads);
  • wt,jIw^I_{t,j} is a weight that token tt gives to indexer head jj (how much to trust that head for this token);
  • ReLU⁡(x)=max⁡(0,x)\operatorname{ReLU}(x) = \max(0, x).

The paper explains both choices: ReLU was chosen "for throughput consideration", and the indexer "can be implemented in FP8".

Piece 2: top-k token selection

For each query token tt, keep only the kk earlier tokens with the highest index scores, and run the real attention over just those:

St=Top-k⁡(It,:),ut=Attn⁡(ht,{ cs:s∈St })S_t = \operatorname{Top\text{-}k}\big(I_{t,:}\big), \qquad u_t = \operatorname{Attn}\big(h_t, \{\, c_s : s \in S_t \,\}\big)

where:

  • It,:I_{t,:} means all of token tt's index scores;
  • StS_t is the set of selected positions;
  • hth_t is the current token's hidden vector;
  • csc_s is the cached key-value entry of token ss (in DeepSeek, the MLA latent from Part 2);
  • utu_t is the attention output.

DeepSeek-V3.2 uses k=2,048k = 2{,}048.

What it saves

So the expensive part (the big attention with 128 heads) becomes linear in length. The cheap part (the indexer) is still quadratic, but it is so small and so fast that it costs much less.

How much less, exactly? DeepSeek-V3.2's published configuration gives the sizes: the indexer has HI=64H^I = 64 heads of size dI=128d^I = 128; the main attention has 128 heads, each scoring against a cached entry of 512+64=576512 + 64 = 576 numbers. Counting multiply-adds for the scoring step only:

dense≈P×128×576,DSA≈P×64×128⏟indexer, FP8+L k×128×576⏟main, top-k\text{dense} \approx P \times 128 \times 576, \qquad \text{DSA} \approx \underbrace{P \times 64 \times 128}_{\text{indexer, FP8}} + \underbrace{L\,k \times 128 \times 576}_{\text{main, top-}k}

where P=L(L+1)/2P = L(L+1)/2 is the number of (query, earlier token) pairs and k=2,048k = 2{,}048. At L=131,072L = 131{,}072 tokens:

Pairs scoredMultiply-adds per pairTotal
Dense MLA8.59 billion73,7286.3×10146.3 \times 10^{14}
DSA indexer8.59 billion8,192 (in FP8)7.0×10137.0 \times 10^{13}
DSA main0.27 billion (at most)73,7282.0×10132.0 \times 10^{13}
DSA total9.0×10139.0 \times 10^{13}, about 7× less

The indexer still touches every pair, but each touch is 9 times cheaper and runs in 8-bit numbers. The main attention touches 32 times fewer pairs. (This counts only the scoring arithmetic. Real speed also depends on memory reads and GPU code; the paper's Figure 3 below shows the measured result.)

How the indexer learns: two training stages

The indexer starts out random. How does it learn which tokens matter? It copies the model's own attention. DeepSeek trains it in two stages:

Stage 1, dense warm-up. Keep normal full attention, freeze the whole model, and train only the indexer to predict where the model's attention goes.

As an equation, with At,s(h)A^{(h)}_{t,s} the main attention weight of head hh from token tt to token ss:

pt,s=∑hAt,s(h)∑s′∑hAt,s′(h)p_{t,s} = \frac{\sum_{h} A^{(h)}_{t,s}}{\sum_{s'} \sum_{h} A^{(h)}_{t,s'}}

The top sums over heads; the bottom divides by the row total so the numbers add up to 1.

The training loss compares the indexer's scores (turned into a distribution by softmax) with that target:

LI=∑tDKL(pt,: ∥ Softmax⁡(It,:))\mathcal{L}^I = \sum_t D_{\mathrm{KL}}\Big(p_{t,:} \,\Big\|\, \operatorname{Softmax}(I_{t,:})\Big)

This stage is short: 1,000 steps, 2.1 billion tokens, learning rate 10−310^{-3}.

Stage 2, sparse training. Switch on top-k selection and train the whole model to work with it: 15,000 steps, 943.7 billion tokens. The indexer keeps learning from the KL loss, but it is trained separately: the paper says "we detach the indexer input from the computational graph for separate optimization", meaning the indexer's learning signal does not flow back into the main model.

9. I trained a lightning indexer on a real model

To see whether this really works, and not just trust the paper, I built a small DSA for Qwen2.5-0.5B and ran stage 1 (the dense warm-up) exactly as described: model frozen, indexer trained with the KL loss against the model's own head-summed, L1-normalised attention.

The indexer, one per layer, 24 in total:

  • 4 indexer heads of size 32, one shared key per token, ReLU, and the weights ww;
  • RoPE on its queries and keys, so it knows positions;
  • a LayerNorm on its input (more on this below);
  • 148,736 parameters per layer, about 8% of the size of the attention layer it serves (1,836,160).
python
class LightningIndexer(torch.nn.Module):
    """I[t, s] = sum_j w[t, j] * ReLU(q[t, j] . k[s])   (DeepSeek-V3.2, eq. 1). One shared key per token."""
    def __init__(self, d_model):
        super().__init__()
        self.norm = torch.nn.LayerNorm(d_model)                        # Qwen's hidden states have a few huge values; tame them
        self.q = torch.nn.Linear(d_model, HI * DI, bias=False)
        self.k = torch.nn.Linear(d_model, DI, bias=False)
        self.w = torch.nn.Linear(d_model, HI, bias=False)

    def forward(self, h):                                              # h: (T, d_model)
        n, h = h.shape[0], self.norm(h)
        q = rope(self.q(h).view(n, HI, DI).transpose(0, 1))            # (HI, T, DI)
        k = rope(self.k(h))                                            # (T, DI)
        return torch.einsum('jt,jts->ts', self.w(h).T, F.relu(q @ k.T))   # (T, T)

The target and the loss, as in the paper (summing the 14 heads, then dividing by the total; the KL loss written as cross-entropy, which differs from KL only by a constant):

python
        A = output[1][0]                                               # (heads, T, T)
        captured[layer]['p'] = (A.sum(0) / A.sum(0).sum(-1, keepdim=True)).detach()   # sum over heads, L1-normalise
python
def kl_loss(ix, h, p):
    I = ix(h).masked_fill(~CAUSAL, float('-inf'))
    logq = torch.log_softmax(I, -1)
    return -(p * logq.masked_fill(~CAUSAL, 0)).sum(-1).mean()          # KL(p || softmax(I)) up to a constant

Training: 400 steps per layer, learning rate 10−310^{-3} (the paper's warm-up rate), on six 2,048-token chunks of The Adventures of Sherlock Holmes. Testing used a different book, Pride and Prejudice, so the indexer never saw the test text.

Result 1: does it pick the tokens the model really attends to?

For each layer and each query token (positions 1,024 to 2,047), I measured what share of the model's real attention lands on the chosen tokens. 1.0 would mean the chosen tokens hold all of the attention.

plain text
budget 64: attention captured  window 0.423  sinks+window 0.721  indexer 0.776  (untrained, layer 0: 0.042)
budget 256: attention captured  window 0.520  sinks+window 0.820  indexer 0.887  (untrained, layer 0: 0.177)
Budget per tokenLast k tokens4 sinks + windowTrained indexerUntrained indexer (layer 1)
6442.3%72.1%77.6%4.2%
25652.0%82.0%88.7%17.7%
  • An untrained indexer catches almost nothing (4.2%), so the training clearly did the work.
  • The trained indexer beats the strong "sinks + window" rule at both budgets. It found the sinks and the recent tokens by itself, plus some important far-away tokens.
0.000.250.500.751.0016121824indexer top-64, 1: 0.63indexer top-64, 2: 0.62indexer top-64, 3: 0.78indexer top-64, 4: 0.83indexer top-64, 5: 0.75indexer top-64, 6: 0.74indexer top-64, 7: 0.81indexer top-64, 8: 0.95indexer top-64, 9: 0.94indexer top-64, 10: 0.91indexer top-64, 11: 0.85indexer top-64, 12: 0.86indexer top-64, 13: 0.82indexer top-64, 14: 0.82indexer top-64, 15: 0.86indexer top-64, 16: 0.84indexer top-64, 17: 0.85indexer top-64, 18: 0.70indexer top-64, 19: 0.89indexer top-64, 20: 0.89indexer top-64, 21: 0.74indexer top-64, 22: 0.71indexer top-64, 23: 0.37indexer top-64, 24: 0.454 sinks + window 60, 1: 0.544 sinks + window 60, 2: 0.584 sinks + window 60, 3: 0.694 sinks + window 60, 4: 0.804 sinks + window 60, 5: 0.724 sinks + window 60, 6: 0.714 sinks + window 60, 7: 0.784 sinks + window 60, 8: 0.934 sinks + window 60, 9: 0.914 sinks + window 60, 10: 0.884 sinks + window 60, 11: 0.804 sinks + window 60, 12: 0.804 sinks + window 60, 13: 0.794 sinks + window 60, 14: 0.764 sinks + window 60, 15: 0.824 sinks + window 60, 16: 0.794 sinks + window 60, 17: 0.794 sinks + window 60, 18: 0.674 sinks + window 60, 19: 0.874 sinks + window 60, 20: 0.874 sinks + window 60, 21: 0.674 sinks + window 60, 22: 0.664 sinks + window 60, 23: 0.174 sinks + window 60, 24: 0.30indexer top-64: 0.454 sinks + window 60: 0.30layershare of attention captured
Layer by layer, the share of real attention caught by 64 tokens. The trained indexer (blue) is above "4 sinks + window" (orange) in every layer. The last two layers spread their attention widely, so 64 tokens catch less there.

Here is one real example: the very last token of a test chunk, in layer 13.

Layer 13, the last token (position 2,047): which earlier tokens matter?top 64 tokens by the model's real attention64 tokens picked by the trained indexer05121,0241,5362,047Both rows agree on 58 of 64 tokens. Many picks are recent tokens (right edge) or the first token (left edge).
For one real query: the 64 tokens the model actually attends to most (top) and the 64 tokens the trained indexer picked (bottom), across positions 0 to 2,047. They agree on 58 of 64. Both include the first token (the sink) and a cluster of recent tokens, plus a few far-away ones.

Result 2: does the model still work with sparse attention?

The real test: run Qwen with every layer using top-k attention chosen by its indexer, and measure perplexity on Pride and Prejudice.

64: last 64 tokens: perplexity 132.7064: last 64 tokens132.7, bar cut off64: 4 sinks + last 60: perplexity 22.4564: 4 sinks + last 6022.464: indexer top-64: perplexity 18.9364: indexer top-6418.9256: last 256 tokens: perplexity 67.90256: last 256 tokens67.9, bar cut off256: 4 sinks + last 252: perplexity 18.83256: 4 sinks + last 25218.8256: indexer top-256: perplexity 17.62256: indexer top-25617.6full attention (2,048): perplexity 17.39full attention (2,048)17.4perplexity on Pride and Prejudice, tokens 1,024 to 2,047 (lower is better)
Perplexity when each token may look at only 64 or 256 earlier tokens, chosen three ways, against full attention (lower is better). The trained indexer comes closest to full attention at both budgets.
Budget per tokenLast k tokens4 sinks + windowTrained indexer, top-kFull attention
64132.722.4518.9317.39
25667.918.8317.6217.39

With 256 tokens per query (full attention reads 1,025 to 2,048 at these positions), the indexer version reaches 17.62, against 17.39 for full attention, without retraining the model at all. With 64 tokens it stays at 18.93, clearly better than the best fixed pattern (22.45).

10. Proof of the runs

Both scripts, exactly as they ran (the same outputs quoted above):

Terminal output of part3_sparse.py: sliding-window checks, receptive field, rolling cache, KV memory and decode timing
Output of part3_sparse.py on an Apple M5 Pro.
Terminal output of part3_qwen.py: window and sink perplexities, indexer warm-up losses per layer, attention captured and sparse perplexities
Output of part3_qwen.py: Qwen2.5-0.5B with windows, sinks and trained lightning indexers. The whole run took 106 seconds.

The impact, and where you meet it

Long context went from a research problem to a standard feature between 2023 and 2025, and the ideas in this part are a big reason why:

  • Mistral 7B (2023) made sliding windows and the rolling cache standard in open models.
  • Gemma 2 and 3, gpt-oss made local and global layers the default way to reach 128K tokens with a manageable cache.
  • StreamingLLM showed how to keep a model running on a never-ending stream; keeping attention sinks is now a standard trick in serving systems.
  • DeepSeek-V3.2 showed that a learned top-k selection can keep quality while cutting long-context cost sharply.

Use cases in one line each:

  • Long documents (contracts, books, codebases): local and global layers keep the cache affordable at 128K tokens.
  • Endless chats and live streams: a rolling window plus attention sinks keeps memory fixed for as long as the stream runs.
  • Reasoning models that write tens of thousands of tokens: sparse attention (DSA) keeps each new token cheap, however long the reasoning gets.
  • Retrieval inside long context ("find the clause that mentions X"): this needs global layers or a learned selector, because a pure window cannot look that far directly.

Which one to use?

Each token looks atKV cacheCan look far back directly?Used by
Full attentionall earlier tokensgrows with lengthyesmost models, in at least some layers
Sliding windowthe last Wfixed at Wno (only indirectly)Mistral 7B
Local + global layerslast W in local layers, all in globallocal layers fixed, global growyes, in global layersGemma 2, Gemma 3, gpt-oss
Window + sinksfirst few + last WfixednoStreamingLLM (inference trick)
DSA (top-k)k tokens chosen by the indexerall kept, but only k read per stepyesDeepSeek-V3.2

Notice that DSA still keeps every token's latent in memory (any of them might be chosen), but reads only kk of them per step. Sliding windows save memory too; DSA saves mainly reading and computing.

All of these still use softmax attention over some set of tokens. Part 4 goes one step further: layers that replace softmax attention with a fixed-size memory that never grows at all (linear attention, Gated DeltaNet), and the hybrid models that mix them with ordinary attention.

Summary

  • Full attention compares every token with every earlier token: T(T+1)/2T(T+1)/2 comparisons per head per layer, 8.6 billion at 128K tokens.
  • Sliding-window attention adds one condition to the mask, 0≤i−j<W0 \le i - j < W. I checked it against PyTorch (4.4×10−164.4 \times 10^{-16}) and showed that tokens outside the window have exactly zero effect.
  • Stacked layers still see far: the reach is L(W−1)+1L(W-1)+1, measured exactly with gradients for 1 to 6 layers.
  • A rolling buffer stores token ii in slot i mod Wi \bmod W. It matched full computation to 1.1×10−151.1 \times 10^{-15}, cut Mistral 7B's cache 8× at 32K, and kept a decode step flat (0.031 ms vs 3.893 ms at 128K).
  • Local + global layers keep some full-attention layers for direct long-range lookups. Gemma 3 27B: 62 GiB → 10.4 GiB at 128K.
  • A model trained with full attention breaks under a plain window (perplexity 17.4 → up to 509) because it loses its attention sink. Keeping 4 first tokens fixes most of it (17.80).
  • DeepSeek Sparse Attention uses a tiny lightning indexer, It,s=∑jwt,jI ReLU⁡(qt,jI⋅ksI)I_{t,s} = \sum_j w^I_{t,j}\,\operatorname{ReLU}(q^I_{t,j} \cdot k^I_s), to pick the top-k tokens, so the main attention costs O(Lk)O(Lk) instead of O(L2)O(L^2).
  • My indexer for Qwen2.5-0.5B, trained only with the warm-up KL loss, caught 77.6% of the real attention with 64 tokens and brought perplexity to 17.62 with 256 tokens (full: 17.39).
  • With DeepSeek-V3.2's real sizes, DSA scores a 128K-token text with about 7× fewer multiply-adds than dense attention, and the paper's measured decode cost at 128K drops from about $2.1 to $0.25 per million tokens.
Run it yourself
  • code/attention/part3_sparse.py: sliding-window checks, receptive field, rolling cache, KV memory, decode timing. Runs in seconds; the timing part uses a GPU if one is available.
  • code/attention/part3_qwen.py: Qwen2.5-0.5B with windows and sinks, plus training and testing the 24 lightning indexers. Downloads the model (about 1 GB) and two public-domain books from Project Gutenberg. About 2 minutes on an Apple M5 Pro.
bash
pip install torch transformers
python part3_sparse.py     # writes results/part3.json
python part3_qwen.py       # writes results/part3_qwen.json

References

  1. A. Q. Jiang et al. Mistral 7B. 2023.
  2. I. Beltagy, M. E. Peters, A. Cohan. Longformer: The Long-Document Transformer. 2020.
  3. R. Child, S. Gray, A. Radford, I. Sutskever. Generating Long Sequences with Sparse Transformers. 2019.
  4. G. Xiao, Y. Tian, B. Chen, S. Han, M. Lewis. Efficient Streaming Language Models with Attention Sinks. ICLR 2024.
  5. Gemma Team. Gemma 2: Improving Open Language Models at a Practical Size. 2024.
  6. Gemma Team. Gemma 3 Technical Report. 2025.
  7. DeepSeek-AI. DeepSeek-V3.2: Pushing the Frontier of Open Large Language Models. 2025.
  8. Model configurations: Mistral-7B-v0.1, Gemma 3 27B, Gemma 2 9B, gpt-oss-20b, Qwen2.5-0.5B.
  9. Texts: Pride and Prejudice and The Adventures of Sherlock Holmes, Project Gutenberg.