Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

29. Sparse attention: windows, top-k and learned indexers

In this chapter

  • When attention can safely read only some keys, and when it can't.
  • Fixed patterns: sliding windows and attention sinks.
  • Content-based selection: top-k attention, and why exact top-k saves nothing by itself.
  • Qwen Sparse Attention: a cheap block indexer that chooses what full attention reads, plus gated attention with partial RoPE, as in Flash-Next.
  • Cache consistency: why selection must give identical results for full, chunked and token-by-token forwards.

You will build

QSAIndexer.forward and GatedAttention.forward in engine/sparse.py: Flash-Next's attention layer.

Time: 5-7 hours. GPU: not needed.

Most attention weights are tiny

A trained model’s attention is usually peaked: for a given query, a handful of keys carry most of the weight and thousands share the rest. If the weight on a key is $10^{-6}$, skipping it changes the output by about $10^{-6}$ of that value. Sparse attention bets on this: read only the keys that matter, and save the memory traffic of the rest.

The bet doesn’t always pay. run.py sparse compares exact top-$k$ attention (keep each query’s $k$ highest-scoring keys, renormalize) with dense attention over 1,024 keys, for flat and for peaked score distributions:

python run.py sparse
{"context": 1024, "score_scale": 1.0, "relative_error": {"budget_16": 4.4,   "budget_64": 1.892, "budget_256": 0.595}}
{"context": 1024, "score_scale": 4.0, "relative_error": {"budget_16": 0.123, "budget_64": 0.026, "budget_256": 0.002}}

With flat scores (random queries and keys), the output is an average over many values, and keeping 256 of 1,024 keys still leaves a 60% error. With peaked scores, 64 keys give 2.6% error and 256 give 0.2%. Real models sit closer to the second case for most heads and most layers, which is why sparse attention works in practice, and why it’s trained into the model rather than applied after the fact.

Fixed patterns: windows and sinks

The simplest selection ignores content. A sliding window lets each query see only the last $w$ keys: memory is $O(w)$ per sequence, and the KV cache becomes a ring buffer. Mistral 7B used $w = 4{,}096$; Gemma alternates windowed and global layers.

Xiao et al. (2023) found a twist: windowed attention collapses when the first tokens leave the window, because models learn to dump excess attention weight on the very first positions (“attention sinks”). Keeping the first few tokens visible forever fixes it.

def sliding_window_allowed(query_positions, key_positions, window, sinks=0):
    """Keys within `window` positions of the query, plus the first `sinks` tokens of the sequence."""
    qp = query_positions[..., :, None]
    kp = key_positions[..., None, :]
    return ((qp - kp) < window) | (kp < sinks)

Your causal_attention takes an allowed mask (Chapter 5), so a window is one more boolean term. The milestone test checks a window of 2 with 1 sink: position 5 sees keys 0, 4 and 5.

Content-based selection: top-k

Windows can’t recall something important from 50,000 tokens ago. Top-k attention selects by score instead:

def topk_attention(q, k, v, budget, query_positions=None):
    """Exact scores, but each query keeps only its `budget` highest-scoring visible keys.
    Teaches selection; it still computes every score, so it saves nothing by itself."""
    t, s = q.shape[-2], k.shape[-2]
    qp = torch.arange(s - t, s, device=q.device) if query_positions is None else query_positions
    scores = (q.float() @ k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
    visible = torch.arange(s, device=q.device)[None, :] <= qp[:, None]
    scores = scores.masked_fill(~visible, float("-inf"))
    keep = scores.topk(min(budget, s), dim=-1).indices
    allowed = torch.zeros_like(scores, dtype=torch.bool).scatter(-1, keep, True) & visible
    return causal_attention(q, k, v, qp, None, allowed)

This version is useful for understanding and for measuring error, but it saves nothing: to find the top $k$ scores it computes all of them, reading every key. The savings come only when something cheaper than attention decides which keys to read. That’s an indexer.

Qwen Sparse Attention

Flash-Next’s 12 attention layers use QSA: a small learned indexer picks, for each query, a budget of 2,048 keys out of the whole context, and full attention reads only those. Three ideas make the indexer cheap:

  1. Blocks, not tokens. The indexer’s keys are averaged over consecutive blocks of 4 tokens (the compression ratio). Selection happens per block: one decision per 4 tokens, and contiguous memory reads.
  2. Small and shared. The indexer has its own tiny projection: 4 query heads of dimension 128, and one key per block shared by all of them, against attention’s 2 KV heads of dimension 256 per token.
  3. A simple score. For query $q$ (with heads $h$) and pooled block key $\bar k_b$: $$ \text{score}(q, b) = \frac{1}{\sqrt{d}} \sum_h \operatorname{ReLU}\big(q_h \cdot \bar k_b\big). $$ The ReLU lets each indexer head vote only for blocks, never against.

Then each query keeps its top $\text{budget}/\text{ratio} = 512$ complete visible blocks, plus the open block, the partially filled block it’s in, which is always readable so recent context is never lost. Both the indexer’s queries and the pooled block keys get RoPE (the block key at the block’s first position), so the indexer is position-aware.

class QSAIndexer(nn.Module):
    """Chooses, for every query, which keys attention may read.  (Your engine: Chapter 29)

    Keys are averaged over consecutive blocks of `ratio` tokens; one shared indexer key per
    block (RMSNorm, then RoPE at the block's first position) is scored against several
    small indexer query heads: score = sum over heads of ReLU(q . k) / sqrt(d). Each query
    keeps its top budget/ratio complete visible blocks plus the incomplete trailing block.
    """

    def __init__(self, hidden, heads, head_dim, budget, ratio, eps=1e-6):
        super().__init__()
        if budget % ratio:
            raise ValueError("budget must be a multiple of the compression ratio")
        self.heads, self.head_dim, self.ratio, self.block_topk = heads, head_dim, ratio, budget // ratio
        self.index_qk_proj = nn.Linear(hidden, (heads + 1) * head_dim, bias=False)
        self.q_layernorm = ZeroCenteredRMSNorm(head_dim, eps)
        self.k_layernorm = ZeroCenteredRMSNorm(head_dim, eps)

    def forward(self, x, positions, rotary_dim, theta, cached_keys=None):
        """x [B, T, D] at absolute positions [T] -> (allowed [B, 1, T, S], all raw keys [B, S, d]).  (Your engine: Chapter 29)"""
        b, t, _ = x.shape
        q, raw = self.index_qk_proj(x).split([self.heads * self.head_dim, self.head_dim], dim=-1)
        q = self.q_layernorm(q.view(b, t, self.heads, self.head_dim))
        cos, sin = rope_cos_sin(positions, rotary_dim, theta, q.dtype)       # [1, 1, T, r]
        q = apply_rope(q, cos.transpose(1, 2), sin.transpose(1, 2))          # broadcast over heads
        keys = raw if cached_keys is None else torch.cat((cached_keys, raw), dim=1)
        s = keys.shape[1]
        n_blocks = s // self.ratio
        key_pos = torch.arange(s, device=x.device)
        visible_tail_start = ((positions + 1) // self.ratio) * self.ratio    # first token of the open block
        tail = (key_pos[None, :] >= visible_tail_start[:, None]) & (key_pos[None, :] <= positions[:, None])
        if n_blocks == 0:
            return tail[None, None].expand(b, 1, t, s), keys
        pooled = keys[:, :n_blocks * self.ratio].view(b, n_blocks, self.ratio, -1).float().mean(2).to(keys.dtype)
        pooled = self.k_layernorm(pooled)
        starts = torch.arange(n_blocks, device=x.device) * self.ratio
        bcos, bsin = rope_cos_sin(starts, rotary_dim, theta, pooled.dtype)   # [1, 1, n, r]
        pooled = apply_rope(pooled, bcos[:, 0], bsin[:, 0])
        scores = torch.relu(torch.einsum("bthd,bnd->btnh", q.float(), pooled.float())).sum(-1) / math.sqrt(self.head_dim)
        # Queries that see the same number of complete blocks run topk over exactly those blocks.
        # (topk breaks ties differently for different vector lengths; ReLU makes ties at 0 common,
        # so slicing rather than masking with -inf is what keeps us identical to the reference.)
        visible_blocks = (positions + 1) // self.ratio                         # [T]
        chosen = torch.zeros(b, t, n_blocks, dtype=torch.bool, device=x.device)
        for count in visible_blocks.unique().tolist():
            if count == 0:
                continue
            rows = visible_blocks == count
            top = scores[:, rows, :count].topk(min(self.block_topk, count), dim=-1).indices
            chosen[:, rows] = chosen[:, rows].scatter(-1, top, True)
        block_of_key = (key_pos // self.ratio).clamp(max=n_blocks - 1)        # tail keys are masked below
        from_blocks = chosen.gather(-1, block_of_key.expand(b, t, s))
        from_blocks = from_blocks & (key_pos < n_blocks * self.ratio)
        return (from_blocks | tail[None])[:, None], keys

The indexer’s output is just an allowed mask [B, 1, T, S], and attention uses it exactly like a causal mask. A real kernel instead gathers the selected blocks’ K and V (Chapter 25’s block tables are the natural fit) and never reads the others.

Explore: which blocks does a query read?

A grid of queries by key blocks. Change the budget, block size and pattern (window, window + sinks, top-k by score) and see which keys each query may read and how many bytes that is.

What it saves

The arithmetic for one Flash-Next attention layer at decode, in BF16 (the second half of run.py sparse):

{"context": 8192,   "dense_kv_MB_per_layer_step": 16.8,  "qsa_MB_per_layer_step": 4.7,  "reduction": 3.5}
{"context": 32768,  "dense_kv_MB_per_layer_step": 67.1,  "qsa_MB_per_layer_step": 6.3,  "reduction": 10.7}
{"context": 262144, "dense_kv_MB_per_layer_step": 536.9, "qsa_MB_per_layer_step": 21.0, "reduction": 25.6}

QSA reads the K and V of 2,048 selected tokens plus every block’s 128-dimensional indexer key. The indexer’s reads still grow with context, but 16 times more slowly than the full KV cache (one 256-byte key per 4 tokens, against 4 KiB of K and V per token). The cache itself still stores everything: sparse attention saves bandwidth and compute, not memory. That’s why Flash-Next combines it with linear attention (Chapter 28), which saves memory.

A subtlety: a query may not see itself

A query at position 11, with blocks of 4, completes block 2. Its open block is now empty, and block 2 is a complete block that must compete in the top-$k$ with blocks 0 and 1. With a budget of two blocks and a low score for block 2, the query attends to blocks 0 and 1 and not to itself. The milestone test asserts that this can happen, rather than assuming that every query sees its own token. It always sees some keys (at least one complete block wins the top-$k$), just not necessarily its own.

Gated attention with partial RoPE

Flash-Next’s attention layer (inherited from Qwen3-Next and Qwen3.5) changes three more things from Chapter 17’s Qwen3:

  • An output gate. q_proj produces twice the query width: a query and a gate per head. The attention output is multiplied by $\sigma(\text{gate})$ before o_proj. The gate lets a head output nothing when it has nothing useful to add, which also removes the need for attention sinks (Qiu et al., 2025, Gated Attention for Large Language Models).
  • Partial RoPE. Only the first 25% of each head’s 256 dimensions are rotated (partial_rotary_factor); the remaining 192 carry content without position. Your apply_rope from Chapter 17 rotates the first rotary_dim features.
  • Zero-centered RMSNorm for the Q and K norms: the scale is $1 + w$ with $w$ initialized at 0, so weight decay pulls toward the identity instead of toward zero.
class ZeroCenteredRMSNorm(nn.Module):
    """RMSNorm with scale (1 + weight), weight initialized at 0. Optionally normalizes each
    group of `group_size` features separately (used on the 4 residual streams in Chapter 30)."""
    def __init__(self, width, eps=1e-6, group_size=None):
        super().__init__()
        self.weight = nn.Parameter(torch.zeros(width))
        self.eps, self.group_size = eps, group_size

    def forward(self, x):
        x32 = x.float()
        if self.group_size:
            x32 = x32.unflatten(-1, (-1, self.group_size))
        x32 = x32 * torch.rsqrt(x32.square().mean(-1, keepdim=True) + self.eps)
        if self.group_size:
            x32 = x32.flatten(-2)
        return (x32 * (1.0 + self.weight.float())).type_as(x)
class GatedAttention(nn.Module):
    """Qwen3.5/Flash-Next attention: zero-centered Q/K norms, partial RoPE, GQA, an optional QSA
    indexer, and a sigmoid output gate produced by the same q_proj.  (Your engine: Chapter 29)"""

    def __init__(self, hidden, heads, kv_heads, head_dim, rotary_dim, theta, eps=1e-6, indexer=None):
        super().__init__()
        self.heads, self.kv_heads, self.head_dim = heads, kv_heads, head_dim
        self.rotary_dim, self.theta = rotary_dim, theta
        self.q_proj = nn.Linear(hidden, heads * head_dim * 2, bias=False)    # [query | gate] per head
        self.k_proj = nn.Linear(hidden, kv_heads * head_dim, bias=False)
        self.v_proj = nn.Linear(hidden, kv_heads * head_dim, bias=False)
        self.o_proj = nn.Linear(heads * head_dim, hidden, bias=False)
        self.q_norm = ZeroCenteredRMSNorm(head_dim, eps)
        self.k_norm = ZeroCenteredRMSNorm(head_dim, eps)
        self.indexer = indexer

    def forward(self, x, positions, state=None):
        """x [B, T, D], positions [T] -> (y [B, T, D], new AttentionState).  (Your engine: Chapter 29)"""
        b, t, _ = x.shape
        state = state or AttentionState()
        query, gate = self.q_proj(x).view(b, t, self.heads, 2 * self.head_dim).chunk(2, dim=-1)
        q = self.q_norm(query).transpose(1, 2)
        k = self.k_norm(self.k_proj(x).view(b, t, self.kv_heads, self.head_dim)).transpose(1, 2)
        v = self.v_proj(x).view(b, t, self.kv_heads, self.head_dim).transpose(1, 2)
        cos, sin = rope_cos_sin(positions, self.rotary_dim, self.theta, q.dtype)
        q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)      # only the first rotary_dim features
        allowed, index_keys = None, None
        if self.indexer is not None:
            allowed, index_keys = self.indexer(x, positions, self.rotary_dim, self.theta, state.index_keys)
        state = state.append(k, v)
        state.index_keys = index_keys                                # all raw indexer keys so far
        y = causal_attention(q, state.k, state.v, positions, None, allowed)
        y = y.transpose(1, 2).reshape(b, t, -1) * torch.sigmoid(gate.reshape(b, t, -1))
        return self.o_proj(y), state

AttentionState carries the layer’s per-request cache: K, V and all raw indexer keys (pooling happens on the fly, so a decode step that completes a block sees it immediately). This reference concatenates; in Chapter 30’s engine this is the state that a paged or static cache would replace.

Consistency: the hard requirement

Selection is discrete: a block is in or out. If a full forward and a token-by-token decode compute scores with tiny differences, they can select different blocks and produce very different outputs, and the cache-equivalence test from Chapter 16 fails in a way that looks random. Two rules keep them identical:

  1. Score each query against exactly the blocks it can see, the same way in both paths. A decode step at position $p$ must see the same pooled keys a full forward does for row $p$.
  2. Break ties the same way. ReLU makes exact ties at 0 common (many blocks score 0). PyTorch’s topk breaks ties differently depending on the length of the vector it searches. Masking unseen blocks with $-\infty$ in a long vector and slicing a short vector to just the visible blocks can choose different zero-scored blocks. The reference therefore groups queries by their number of visible blocks and runs topk over exactly that slice, for every path.

The second rule was found the hard way: the first version of the whole Flash-Next model differed from Hugging Face’s on a few positions, and the cause was exactly this tie-breaking. The milestone test runs prefill of 11 tokens followed by 14 single-token steps against one full forward, with a budget small enough that selection matters.

Build it

Engine milestone 29: sparse attention. Implement QSAIndexer.forward and GatedAttention.forward in engine/sparse.py (windows, the norm, top-k attention and the state are provided).

pytest tests/test_ch29_sparse.py
python run.py sparse --impl engine

The tests check windows with sinks, that top-k with a full budget equals dense attention, the grouped zero-centered norm, that the indexer always keeps the open block, respects the budget and never selects a future key, and that prefill plus token-by-token decode with selection equals one full forward. Chapter 30’s parity test checks the layer against Hugging Face’s implementation inside the full model.

Stretch exercises

  1. ★ Add an attention-sink option to topk_attention (always keep the first 4 keys) and measure the error change with flat and peaked scores. Where: topk_attention in engine/sparse.py.
  2. ★★ Implement a ring-buffer KV cache for sliding-window layers, storing only the last $w$ positions, and verify a windowed model’s logits against the full-cache version. Where: AttentionState in engine/sparse.py, with position handling in GatedAttention.forward.
  3. ★★ Make the indexer’s selection explicit: return the selected block indices [B, T, budget/ratio] instead of a mask, and write attention that gathers only those blocks’ K and V. Check it against the mask version. Where: QSAIndexer.forward and GatedAttention.forward in engine/sparse.py.
  4. ★★★ Write a Triton decode kernel for QSA: one program per (sequence, head) that reads the selected block list and streams those blocks’ K/V with the online softmax (Chapter 25’s paged kernel is most of it). Where: new engine/kernels/triton_sparse.py, called by GatedAttention.forward in engine/sparse.py.

Check your understanding

  1. Why does top-k attention with exact scores save no computation?
  2. Why does Flash-Next’s indexer select blocks rather than tokens?
  3. What is the open block, and why is it always readable?
  4. Why does sparse attention save bandwidth but not cache memory?
  5. Why must full and incremental forwards break top-k ties identically?

Going deeper

  • Beltagy et al., Longformer (2020) and Zaheer et al., BigBird (2020) for fixed sparse patterns; Xiao et al., Efficient Streaming Language Models with Attention Sinks (2023).
  • DeepSeek-AI, Native Sparse Attention (2025) and the DeepSeek-V3.2 report (DeepSeek Sparse Attention, a lightning indexer with ReLU scores): the closest published relatives of QSA.
  • Qiu et al., Gated Attention for Large Language Models: Non-linearity, Sparsity, and Attention-Sink-Free (2025), the gate used here.
  • The Transformers qwen4_exp modeling file, the reference implementation this chapter’s code is checked against.