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.
In full attention, token number i compares its query with the keys of tokens 1,2,…,i. So for a text of T tokens, the number of query-key comparisons in one head of one layer is:
1+2+3+⋯+T=2T(T+1)
Here is what that means in numbers:
Text length T
Comparisons, full attention
Comparisons, window of 1,024
Full ÷ window
1,024 tokens
524,800
524,800
1.0×
8,192 tokens
33,558,528
7,864,832
4.3×
131,072 tokens (128K)
8,590,000,128
133,693,952
64.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 d (about 2d multiply-and-add steps), and the weighted mix of values costs about the same again. So per head and layer:
where T is the text length, W the window and d the head size. The first grows with T2, the second only with T: 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.
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).
I checked three things in part3_sparse.py (24 tokens, W=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
My function agrees with PyTorch's built-in attention to 4.4×10−16 (rounding noise).
A window wider than the text is exactly ordinary causal attention (difference 0).
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.
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−1 more tokens back. After L layers:
reach=L×(W−1)+1tokens
The Mistral 7B paper says the same thing (counting the window with its W+1 convention):
Longformer (2020) wrote the same rule in its own notation, with ℓ layers and window w:
Proof. I stacked 1 to 6 sliding-window layers (W=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 W steps. So the KV cache can be a fixed-size ring.
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 W slots:
python
k_cache = torch.zeros(H, W, d, dtype=torch.float64) # the whole cache: W slots, never morev_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:
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.
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.
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:
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 far
Full attention
Window of 1,024
1,024
0.057 ms
0.041 ms
4,096
0.137 ms
0.033 ms
16,384
0.521 ms
0.030 ms
65,536
2.019 ms
0.034 ms
131,072
3.893 ms
0.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.)
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:
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:
Global layers store every token; local layers store at most W:
KV cache=2Hkvdhb(Lglobal⋅n+Llocal⋅min(n,W))
where Lglobal and Llocal are the numbers of global and local layers. For Gemma 3 27B (Hkv=16, dh=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)
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 i of a head of size d is
θi=base−2i/d,i=0,1,…,2d−1
where "base" is 10,000 normally and 1,000,000 in Gemma 3's global layers. The slowest pair (i=d/2−1) has a period, in tokens, of about 2π⋅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.
Model
Pattern
Window
Why it matters
Mistral 7B
every layer local
4,096
made SWA and the rolling cache popular in open models
Gemma 2
1 local : 1 global
4,096
half the layers stay small
Gemma 3
5 local : 1 global
1,024
6× smaller cache at 128K
gpt-oss
1 local : 1 global
128
tiny 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.39budget 64 tokens: window only 132.70 4 sinks + window 60: 22.45budget 256 tokens: window only 67.90 4 sinks + window 252: 18.83budget 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 S sink tokens and a window of W recent tokens, the cache holds a fixed
cache size=(S+W)×bytes per token
whatever the length of the conversation: 4 + 1,020 = 1,024 tokens in my 1,024-budget test.
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.
DSA arrived with DeepSeek-V3.2 (2025). It has two pieces.
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.
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=64 heads of size dI=128; the main attention has 128 heads, each scoring against a cached entry of 512+64=576 numbers. Counting multiply-adds for the scoring step only:
where P=L(L+1)/2 is the number of (query, earlier token) pairs and k=2,048. At L=131,072 tokens:
Pairs scored
Multiply-adds per pair
Total
Dense MLA
8.59 billion
73,728
6.3×1014
DSA indexer
8.59 billion
8,192 (in FP8)
7.0×1013
DSA main
0.27 billion (at most)
73,728
2.0×1013
DSA total
9.0×1013, 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.)
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) the main attention weight of head h from token t to token s:
pt,s=∑s′∑hAt,s′(h)∑hAt,s(h)
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=t∑DKL(pt,:Softmax(It,:))
This stage is short: 1,000 steps, 2.1 billion tokens, learning rate 10−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.
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 w;
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−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.
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.
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.
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.
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 token
Last k tokens
4 sinks + window
Trained indexer, top-k
Full attention
64
132.7
22.45
18.93
17.39
256
67.9
18.83
17.62
17.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).
Both scripts, exactly as they ran (the same outputs quoted above):
Output of part3_sparse.py on an Apple M5 Pro.Output of part3_qwen.py: Qwen2.5-0.5B with windows, sinks and trained lightning indexers. The whole run took 106 seconds.
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.
Notice that DSA still keeps every token's latent in memory (any of them might be chosen), but reads only k 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.
Full attention compares every token with every earlier token: T(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<W. I checked it against PyTorch (4.4×10−16) and showed that tokens outside the window have exactly zero effect.
Stacked layers still see far: the reach is L(W−1)+1, measured exactly with gradients for 1 to 6 layers.
A rolling buffer stores token i in slot imodW. It matched full computation to 1.1×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,jIReLU(qt,jI⋅ksI), to pick the top-k tokens, so the main attention costs O(Lk) instead of O(L2).
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.