Fast Attention and Long Context

Sliding-window attention

Sliding-window attention lets each token look only at a fixed number of recent tokens, which caps memory forever, while stacked layers still pass information across the whole document.

On this page 9
  1. The short answer
  2. The line of people passing a message
  3. Why anyone would restrict the model
  4. How far can it really see?
  5. What real models do
  6. The honest limit
  7. Where you have already seen this
  8. Remember this
  9. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

Sliding-window attention lets each word look back at only the last few thousand words, not all of them.

The line of people passing a message

Twenty people stand in a line. Each one may whisper only to the person beside them.

Tell the first person a sentence. It travels down the line, one hop at a time, and reaches the last person. Nobody spoke across the room, and the message still crossed it.

That is a stack of sliding-window layers. Each layer looks only at nearby words. Information still travels far, by being passed along.

And you already know the catch from playing this game. What arrives at the end is not always what was said at the start.

Why anyone would restrict the model

Ordinary attention lets every word look at every earlier word. Cost grows with the square of the length. Doubling the document quadruples the work.

Worse, the model must store notes on every word it has read. Those notes are the KV cache, and it grows forever as the conversation grows.

A window fixes both. Look back at most a fixed number of words. Store at most that many notes. The cost of the ten-thousandth word equals the cost of the hundredth.

   full attention                sliding window (last 4)
   --------------                -----------------------
   t5 sees t0 t1 t2 t3 t4 t5     t5 sees       t2 t3 t4 t5
   t6 sees t0 t1 t2 t3 t4 t5 t6  t6 sees          t3 t4 t5 t6
   t7 sees everything before it  t7 sees             t4 t5 t6 t7

   notes grow forever            notes stop growing

How far can it really see?

One layer with a window of four sees four words back. Two layers see eight, because the second layer reads words that already absorbed their own neighbourhood.

Stack thirty-two layers with a window of four thousand and the reach is over a hundred thousand words. Which is a real reach, and a weaker one than direct attention.

Direct attention is a phone call. Passing along the line is a chain of whispers. Both deliver the message. Only one delivers it exactly.

What real models do

Most current models do not use windows everywhere. They mix.

Some layers use a small window and stay cheap. A few layers keep full attention and can reach anything directly. The cheap layers do the bulk of the work, and the expensive layers rescue the long-range cases.

One popular model alternates them, one for one. Another uses five windowed layers for every full one. The mixture is tuned, not obvious.

The honest limit

A window is not free quality. Suppose a model must recall an exact phrase from very early in a long document. Narrow windows make it measurably worse at that.

If your task is "find this one line in a hundred pages", windowing is working against you. If your task is "write the next paragraph of this chapter", it costs you almost nothing.

Where you have already seen this

  • Local models that keep a steady speed deep into a long chat.
  • Phone keyboards predicting your next word from the current sentence, not your life story.
  • Long-document tools that summarise section by section rather than all at once.

Remember this

  • Each word attends to a fixed number of recent words. Cost per word stops growing.
  • Stacked layers still reach far, indirectly, like a message down a line.
  • Most models mix a few full-attention layers in, because the indirect path is not as reliable.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written against PyTorch 2.5.1, Python 3.10. CPU only, runs instantly.

The mask, the reach, and the cache

sliding_window.py
import torch, torch.nn.functional as F

torch.manual_seed(0)
N, W = 12, 4                      # 12 tokens, each looks back at most 4 (itself + 3)

i = torch.arange(N)[:, None]
j = torch.arange(N)[None, :]
causal = j <= i
window = causal & (i - j < W)     # the sliding-window mask

print("allowed (1) vs blocked (.) - rows are queries, columns are keys")
for r in range(N):
    print(f"  t{r:<2}", " ".join("1" if window[r, c] else "." for c in range(N)))

print()
print("attended per query:", window.sum(1).tolist())
print("full causal would be:", causal.sum(1).tolist())

# What one layer can see directly, vs what L layers can see through each other.
reach = window.clone()
print()
print(f"{'layers':>7}{'oldest token t11 can reach':>28}")
for L in range(1, 5):
    oldest = int(reach[N - 1].nonzero().min())
    print(f"{L:>7}{'t' + str(oldest):>28}")
    reach = (reach.float() @ window.float()) > 0     # compose one more layer

# Attention still runs: feed the mask straight to PyTorch.
q, k, v = (torch.randn(1, 2, N, 8) for _ in range(3))
out = F.scaled_dot_product_attention(q, k, v, attn_mask=window)
print()
print("output shape with the window mask:", tuple(out.shape))

print()
print("KV cache per token, bytes at bf16, 128k context vs a bounded window")
def cache(L, kvh, dh, ctx, win=None):
    eff = ctx if win is None else min(ctx, win)
    return 2 * L * kvh * dh * 2 * eff
rows = [("Mistral-7B-v0.1, window 4096", 32, 8, 128, 4096),
        ("Mistral-7B-v0.1, no window", 32, 8, 128, None),
        ("gpt-oss-20b, 12 sliding layers", 12, 8, 64, 128),
        ("gpt-oss-20b, 12 full layers", 12, 8, 64, None)]
for name, L, kvh, dh, win in rows:
    full = cache(L, kvh, dh, 131072)
    got  = cache(L, kvh, dh, 131072, win)
    print(f"  {name:<28} {got/2**30:7.3f} GB   (unbounded would be {full/2**30:6.2f} GB)")
Output
allowed (1) vs blocked (.) - rows are queries, columns are keys
  t0  1 . . . . . . . . . . .
  t1  1 1 . . . . . . . . . .
  t2  1 1 1 . . . . . . . . .
  t3  1 1 1 1 . . . . . . . .
  t4  . 1 1 1 1 . . . . . . .
  t5  . . 1 1 1 1 . . . . . .
  t6  . . . 1 1 1 1 . . . . .
  t7  . . . . 1 1 1 1 . . . .
  t8  . . . . . 1 1 1 1 . . .
  t9  . . . . . . 1 1 1 1 . .
  t10 . . . . . . . 1 1 1 1 .
  t11 . . . . . . . . 1 1 1 1

attended per query: [1, 2, 3, 4, 4, 4, 4, 4, 4, 4, 4, 4]
full causal would be: [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]

 layers  oldest token t11 can reach
      1                          t8
      2                          t5
      3                          t2
      4                          t0

output shape with the window mask: (1, 2, 12, 8)

KV cache per token, bytes at bf16, 128k context vs a bounded window
  Mistral-7B-v0.1, window 4096   0.500 GB   (unbounded would be  16.00 GB)
  Mistral-7B-v0.1, no window    16.000 GB   (unbounded would be  16.00 GB)
  gpt-oss-20b, 12 sliding layers   0.003 GB   (unbounded would be   3.00 GB)
  gpt-oss-20b, 12 full layers    3.000 GB   (unbounded would be   3.00 GB)

Reading the output

The mask is a band, not a triangle. Compare attended per query: it climbs to 4 and stops, while full causal attention climbs to 12. That flat tail is the property that makes cost per token constant.

The reach table is the receptive field. One layer reaches t8, four layers reach t0. Each extra layer buys another W - 1 tokens of reach. With L layers and window W, reach is about L x (W - 1).

That composition is done with a boolean matrix product, which is a reachability computation on the attention graph. It tells you what can influence a token, not how strongly it does.

The cache table is the reason for the whole design. Mistral-7B's window turns a 16 GB cache at 128k context into 0.5 GB. gpt-oss-20b's twelve windowed layers hold 3 MB at any context length, because 128 tokens is 128 tokens forever.

Note the asymmetry in that last model: 12 full-attention layers hold 3 GB, the 12 windowed layers hold 0.003 GB. Once you mix, the full layers are the entire budget.

What real models set

Verified from each model's config.json on the Hugging Face Hub:

ModelWindowPattern
Mistral-7B-v0.14096every layer
gpt-oss-20b128alternating: sliding, full, sliding, full, … over 24 layers
Longformer-base512every layer, plus designated global tokens
Gemma 31024five local layers per one global layer

Read gpt-oss-20b's layer_types field directly and you get the literal list, ['sliding_attention', 'full_attention', …]. That field is the clearest documentation of a hybrid design that exists.

Google reports that Gemma 3's 5:1 ratio with a 1024 window cut long-context KV cache overhead from around 60% of memory to under 15%.

Doing it properly in a kernel

The boolean mask above is a teaching device. A dense N x N mask allocates the exact tensor that FlashAttention exists to avoid.

Production paths:

  • torch.nn.attention.flex_attention with a mask_mod closure compiles the window into a fused kernel and skips whole blocks that are entirely masked. This is a prototype API as of PyTorch 2.13, so pin your version.
  • flash-attn accepts a window_size=(left, right) tuple directly.
  • vLLM and llama.cpp implement windowed cache eviction, so the cache buffer itself never grows past the window.

The performance win only arrives when blocks are skipped. A dense mask gives you the correct answer at full cost.

Common mistakes

Setting a window and still allocating a full cache. The memory saving comes from evicting old entries, not from masking them. Check your server actually rolls the buffer.

Assuming a windowed model can retrieve from the far past. It can, through layer composition, and less reliably than a full-attention model. Test on your own retrieval task before promising it.

Off-by-one in the window. i - j < W includes the current token, giving W attended positions. i - j <= W gives W + 1. Different papers and libraries pick different conventions; print the mask.

Forgetting the first tokens. Rows t0 to t3 attend to fewer than W tokens. Combined with softmax this creates the pathology that attention sinks exist to fix, and it is the topic of the next lesson.

Comparing window sizes across models without the layer count. Reach is roughly window times depth. A 128-token window over 24 layers and a 4096-token window over 32 layers are not remotely the same design.

Try it yourself

Set W = 2 and rerun the reach loop with more iterations. Count how many layers are needed to reach t0, and check it against L x (W - 1). Then build a hybrid: apply window on odd layers and causal on even ones, and watch reach jump to full in a single step.

What to learn next

Researcher — Mathematics and papers.

Definition

Replace the causal mask $M_{ij} = \mathbb{1}[j \le i]$ with a banded mask

$$ M_{ij} = \mathbb{1}[\,i - w < j \le i\,] $$

where $w$ is the window width. Attention cost per layer falls from $\Theta(N^2 d)$ to $\Theta(N w d)$, and KV cache from $\Theta(N)$ to $\Theta(\min(N, w))$ per layer. Both become linear in sequence length, with $w$ the constant.

Receptive field

Treat attention as a directed graph on positions. One layer connects $i$ to ${i-w+1, \dots, i}$. After $L$ layers the reachable set from position $i$ extends back approximately

$$ r = L (w - 1) + 1 $$

Mistral-7B: $32 \times 4095 + 1 = 131{,}041$, essentially its advertised context. That equality is not accidental; the window was chosen to make it hold.

Reachability is a necessary condition, not a sufficient one. Luo et al. (2016), Understanding the Effective Receptive Field, showed for CNNs that effective influence decays sharply inside the theoretical field and grows as $O(\sqrt{L})$ rather than $O(L)$. The analogous decay in stacked local attention is why hybrid designs exist at all: if composition were as good as direct attention, no model would keep global layers.

Lineage

  • Child, Gray, Radford and Sutskever (2019), Generating Long Sequences with Sparse Transformers (arXiv:1904.10509) — strided and fixed local patterns, factorised so that two layers cover all positions.
  • Beltagy, Peters and Cohan (2020), Longformer (arXiv:2004.05150) — sliding window plus dilation plus task-specific global tokens. config.json still reports attention_window: [512] * 12.
  • Jiang et al. (2023), Mistral 7B (arXiv:2310.06825) — window 4096 at every layer, in a decoder trained from scratch, with a rolling buffer cache.
  • Gemma Team (2025), Gemma 3 Technical Report (arXiv:2503.19786) — 5:1 local-to-global with $w = 1024$, reported to cut long-context KV overhead from roughly 60% to under 15%.
  • OpenAI (2025), gpt-oss model card (arXiv:2508.10925) — alternating banded and dense layers, $w = 128$, with learned attention sinks.

The trend across these is unambiguous: windows have got narrower while the ratio of local to global layers has got higher. gpt-oss's 128 is two orders of magnitude below Longformer's 512-per-layer-everywhere design, because a handful of full layers now carries the long-range load.

Interaction with position encoding

A windowed layer never sees a relative offset larger than $w$. This bounds the range of rotary frequencies the layer must represent, and is one reason windowed models extrapolate more gracefully than their global-attention counterparts at inference lengths beyond training.

Hybrid architectures often exploit this deliberately: global layers get a large RoPE base (long wavelength, better extrapolation) while local layers keep a small base. Gemma 3 does exactly this.

The failure mode

Softmax is a normaliser. Every query must distribute a total weight of one across its window, whether or not any key in the window deserves it. When the window slides past the tokens a head wanted to attend to, the head is forced to place that weight somewhere.

Empirically models solve this by dumping weight on the first few tokens of the sequence. Evict those from a rolling buffer and perplexity explodes. That observation is the entire subject of the next lesson, attention sinks, and it is the reason a naive rolling-cache implementation of sliding-window attention fails in production despite being mathematically what the paper describes.

When not to use a window

Retrieval-shaped tasks — exact quotation, needle-in-a-haystack lookup, cross-referencing two distant tables — depend on direct long-range edges. Windowed layers degrade on these in a way that aggregate perplexity hides almost completely. Evaluate with RULER-style tasks, not with loss curves.

What to learn next