Standard attention lets every token look at every earlier token. That is what makes transformers good at long-range reasoning, and it is also why their cost grows with the square of the sequence length. Sparse attention is the family of designs that computes only some of those query-key pairs and skips the rest, on purpose, in a way the model can still learn around.

The hard part is not skipping work. It is deciding which work to skip without losing the token the model needed, and then skipping it in a way a GPU actually runs faster. This article covers the whole design space: fixed patterns chosen before training, the block-sparse kernels that make them fast, learned selection methods where the model picks its own keys, and the local/global layer mixes used by recent open models. A worked example at 128K tokens shows what each choice saves in compute, KV-cache memory and decode bandwidth, and what it costs.

Advertisement

Why dense attention becomes the bottleneck

For one head, attention computes a score for every query-key pair, applies a softmax over each row, and uses those weights to mix the value vectors. With sequence length n and head dimension d, the two matrix multiplies cost about 4·n·m·d floating-point operations, where m is the number of keys each query reads. Dense causal attention has m ≈ n/2 on average, so the cost is quadratic in n.

Kernels such as FlashAttention fixed the memory side: they never write the n × n score matrix out to HBM. They did not change the arithmetic. At 4K tokens attention is a modest share of a layer's work next to the MLP; at 128K it dominates prefill. During decoding the problem is bandwidth: every new token reads the keys and values of every earlier token from the KV cache, in every layer.

Sparse attention attacks m directly. If each query reads a fixed budget of keys, say 2,048, then the cost becomes linear in n and the quadratic term disappears, provided the selection step is itself cheap.

The design space: who decides which pairs interact

Every sparse attention method answers one question: for this query, which keys? The answers fall into three groups, and most production systems combine them.

ApproachWho chooses the keysExamplesMain risk
Fixed patternThe architect, before trainingSliding window, strided (Sparse Transformer), global tokens (Longformer), random links (BigBird)The needed token falls outside the pattern
Learned selectionThe model, per query, at run timeRouting and LSH attention, NSA, MoBA, DeepSeek Sparse AttentionSelection misses; selection cost; harder training
Layer mixingThe architect, per layerAlternating local and full-attention layersFull layers still pay quadratic cost and full KV memory
Four ways to decide which query-key pairs get computed (16 tokens, causal; rows = queries, columns = keys)Local window (w=4)Local + strided summariesLocal + global tokensTop-k blocks (learned)computedskipped (sparse)future (causal mask)The first three masks are fixed before training. The fourth depends on the content of each query and is recomputed every step.Kernels skip whole 4x4 tiles here; real kernels skip 64x64 or 128x128 tiles, so only block-aligned patterns save time.
Causal masks for four sparse patterns. Fixed patterns are known before training; learned top-k selection changes with every query block.
Advertisement

Fixed patterns and why reachability matters

A sliding window lets each token attend to the previous w tokens. It works because most of the information a token needs is nearby, and because stacking layers widens the reach: after L layers information can travel about L·w positions. The window arithmetic and the rolling-buffer cache are derived in the sliding-window math article; the important caveat here is that multi-hop reach is not the same as direct access. A fact that must travel 30 hops through intermediate tokens is diluted along the way.

The Sparse Transformer of 2019 added strided connections: besides its local block, each query reads a set of summary positions spaced at a fixed stride, so any two tokens are connected within two hops and cost falls to about n·√n. Longformer added a handful of global tokens that attend everywhere and are attended by everyone, such as a classification token or the question in a QA prompt. BigBird combined windows, global tokens and random links, and used the graph argument explicitly: a sparse graph with a few global hubs and random edges has short paths between all nodes.

Fixed patterns are easy to implement and to reason about, and they cost nothing to select. Their weakness is that they are blind to content. If the answer to a question sits 90,000 tokens back and is neither local nor a designated global token, the model has to route it through many layers or lose it.

Block-sparse kernels: sparsity a GPU can use

A GPU attention kernel processes queries and keys in tiles, typically 64 or 128 rows by 64 or 128 columns, because tensor cores multiply small dense matrices. Skipping individual scores inside a tile saves nothing: the tile is computed anyway and the masked scores are set to minus infinity. Skipping a whole tile saves the full cost of that tile. Practical sparse attention is therefore block sparse: the pattern is expressed at tile granularity, plus an optional element-wise mask for tiles that are only partly kept, such as the diagonal tiles of a causal mask.

PyTorch's FlexAttention exposes exactly this split. You write a mask_mod that says whether query position q_idx may see key position kv_idx; create_block_mask evaluates it once per tile and records which tiles are empty, full or partial, and the compiled kernel skips the empty ones.

import torch
from torch.nn.attention.flex_attention import flex_attention, create_block_mask

WINDOW, N_GLOBAL = 4096, 64          # local window plus 64 leading "global" tokens

def local_plus_global(b, h, q_idx, kv_idx):
    causal = q_idx >= kv_idx
    local = (q_idx - kv_idx) < WINDOW
    is_global = kv_idx < N_GLOBAL     # every query may read the prompt header
    return causal & (local | is_global)

S = 131072
block_mask = create_block_mask(local_plus_global, B=None, H=None, Q_LEN=S, KV_LEN=S)
print(block_mask.sparsity())          # percentage of tiles skipped

attn = torch.compile(flex_attention)
out = attn(q, k, v, block_mask=block_mask)   # q, k, v: [batch, heads, S, head_dim]

Two practical rules follow. First, align windows and global regions to the tile size where you can; a 4,000-token window wastes part of a tile at each edge, while 4,096 does not. Second, causal block-sparse work is unevenly distributed: late query blocks read far more tiles than early ones in a full-attention layer, but in a windowed layer every block reads about the same number. The kernel's scheduler matters, and measured speed-ups are usually smaller than the fraction of tiles skipped.

Learned selection: letting the model pick its keys

Content-based methods compute a cheap estimate of which keys matter for each query and run exact attention only on those. Early attempts clustered queries and keys with locality-sensitive hashing (Reformer) or k-means (Routing Transformer). They saved FLOPs on paper but mapped badly onto GPUs and were rarely used in large models. Three recent designs are built around block-sized or hardware-friendly selection instead.

  • Native Sparse Attention (NSA), from DeepSeek researchers, runs three branches per query and combines them with learned gates: attention over compressed tokens (each block of keys summarised into one coarse token), attention over a few selected blocks chosen using the compressed scores, and a sliding window for local context. Selection is at block granularity so the kernel stays dense inside each block, and the model is trained with sparsity from the start rather than converted afterwards.
  • MoBA (Mixture of Block Attention), from Moonshot AI, applies the mixture-of-experts idea to context: keys are split into blocks, each block is scored against the query using a pooled representative of its keys, and the query attends to its top-k blocks plus its own block. It is designed so a model can switch between full and block-sparse attention.
  • DeepSeek Sparse Attention (DSA), introduced in DeepSeek-V3.2-Exp, selects individual tokens rather than blocks. A small lightning indexer, with few heads, a small dimension and FP8 arithmetic, scores every preceding token for each query; exact attention then runs over only the top 2,048 tokens. All heads share the selected set, which fits DeepSeek's multi-query style of latent attention. The indexer is still quadratic, but with a much smaller constant than the main attention.
# Token-level top-k selection in the style of DSA (illustrative pseudocode).
def sparse_attention_step(x, kv_cache, idx_cache, top_k=2048):
    q_idx = indexer_query(x)                 # tiny: few heads, small dim, low precision
    scores = q_idx @ idx_cache.T             # [n_prev]  cheap but O(n) per token
    keep = topk(scores, min(top_k, len(scores))).indices
    k, v = kv_cache.keys[keep], kv_cache.values[keep]   # gather 2,048 entries
    q = main_query(x)
    return softmax(q @ k.T / sqrt(d)) @ v    # exact attention over the subset

Selection is not differentiable, so these methods need a training signal for the selector. NSA's gates and compressed branch provide gradients end to end. DeepSeek's report describes first training the indexer during a short dense warm-up to imitate the main attention distribution, then continuing training with sparse selection switched on. Expect any converted model to need continued pre-training on long documents, not just a code change.

Local and global layers in production models

The most widely deployed form of sparsity is the simplest: make most layers local and keep a few layers with full attention. Mistral 7B used a 4,096-token sliding window. Gemma 2 alternated local and global layers, and Gemma 3 moved to five local layers for every global one with a 1,024-token window. OpenAI's gpt-oss models alternate full attention with banded local attention of 128 tokens.

The appeal is operational. Local layers keep only w entries in their KV cache, so memory and decode bandwidth fall roughly in proportion to the share of local layers, and any kernel that supports a window can run them. The full layers guarantee that every token remains directly reachable somewhere in the network, which is what long-context retrieval depends on. Ring attention and context parallelism remain necessary for those full layers when one device cannot hold the sequence.

Worked example: 128K tokens, three designs

Take a 32-layer model with 32 query heads of dimension 128 and an input of 131,072 tokens. For dense causal attention, one head in one layer costs about 4 × 131,072 × 65,536 × 128 ≈ 4.4 TFLOP during prefill, or about 4.5 PFLOP across all heads and layers. Compare three alternatives on prefill compute, decode reads per new token, and KV-cache size, all relative to dense.

DesignPrefill attention FLOPsKV read per decoded tokenKV cache size
Dense, every layer1×1×1×
5 local (w = 1,024) : 1 full(5/64 + 1)/6 ≈ 0.18×(5 × 1,024 + 131,072)/(6 × 131,072) ≈ 0.17×≈ 0.17×
Top-2,048 tokens per query, every layer≈ 1/32 of main attention, plus the indexer2,048 entries plus indexer keys1×, plus the indexer cache

The last column is the one teams miss. Learned top-k selection does not let you throw tokens away, because any of them might be selected by a later query. It reduces compute and bandwidth, not cache capacity, and it adds a small extra cache for the indexer keys. Local layers reduce all three, but only for the local layers. If KV memory is your constraint, combine sparsity with a smaller cache format, as described in the KV cache article.

Failure modes

  • Silent retrieval loss. Perplexity barely moves while needle-in-a-haystack, multi-document QA and long-code tasks regress, because most next-token predictions only need local context. Evaluate on long-range tasks at the target length.
  • Lost attention sinks. Models dump spare attention mass on the first tokens. A pattern that drops those tokens destabilises the softmax; keep them as global tokens, as explained in attention sinks.
  • Empty rows. A query whose mask keeps no keys produces a softmax over nothing, which becomes NaN. Always keep at least the diagonal.
  • Selection misses at block granularity. One crucial token inside an otherwise irrelevant block scores low after pooling. Smaller blocks and a sliding-window branch reduce the risk at a cost in kernel efficiency.
  • No speed-up at short lengths. Below a few thousand tokens the selector, mask construction and partial tiles can cost more than they save. Many systems fall back to dense attention for short sequences.
  • Train/serve mismatch. A model trained with one window or top-k budget and served with another behaves unpredictably. Pin these values in the model config and assert them in the serving stack.

Operational guidance and trade-offs

Start from the constraint. If prefill compute at long context is the problem, any sparse pattern helps. If decode bandwidth is the problem, local layers and top-k selection both help. If KV memory is the problem, only local layers, cache compression or fewer KV heads help. If retrieval quality is the problem, do not remove full attention from every layer.

Prefer a pattern your inference engine already supports with a fused kernel; a theoretically better pattern that falls back to a masked dense kernel is slower than dense attention. When converting a dense model, budget for continued pre-training and keep a dense baseline to compare against on the same long-context evaluation suite. Measure speed as wall-clock prefill time and decode tokens per second at your real batch sizes, not as the fraction of tiles skipped.

What to do next

  1. Profile one long request and record the attention share of prefill time and the KV-cache size per sequence at your target length.
  2. Write the candidate pattern as a FlexAttention mask_mod, print its block sparsity, and benchmark it against dense FlashAttention at 8K, 32K and 128K.
  3. Choose tile-aligned window sizes and always keep the first few tokens and the diagonal.
  4. Build an evaluation set with needle retrieval, multi-document QA and long-code tasks at the target length; run it before and after any change.
  5. If converting a dense model, plan a dense warm-up for any selector and continued training on long documents.
  6. Pin window size, top-k and layer pattern in the model config and check them at serving start-up.
Key takeaway: Sparse attention trades guaranteed all-pairs access for linear cost. Fixed patterns (windows, strides, global tokens) cost nothing to select but are blind to content; learned selection such as NSA, MoBA and DSA picks keys per query but needs a trained selector and keeps the full KV cache; local/global layer mixes are the simplest way to cut compute, memory and bandwidth together. Make the pattern block-aligned so kernels can skip whole tiles, keep sinks and some full attention for retrieval, and judge every change on long-range tasks rather than perplexity.