Get Started
Home
Topics
Search
Library
7 min read · Inference Optimization · LLM Training · Added Sep 29 · Paper published Sep 25, 2026

Block Sparse Attention with Log-Linear Complexity

Source: research paper via Hugging Face Daily Papers
0:00 / 7:52
PISA attacks the hidden bottleneck in block-sparse attention: picking which blocks to attend to is still O(N²/C). A coarse-to-fine pyramid search with logsumexp scoring drops selection to O(N log N), yielding ~10× speedup at 256K tokens while matching baseline retrieval quality.
TL;DR
PISA replaces the quadratic scan that block-sparse attention uses to pick which key blocks each query reads, doing a coarse-to-fine pyramid search with LogSumExp scoring so selection costs O(N log N) instead of O(N²), while matching baseline quality.
Why It Matters
Long-context transformers are expensive because attention compares every query to every key. Block-sparse attention is a common fix: chop the keys into blocks of, say, 64 tokens, pick the top-K most relevant blocks per query, and only compute attention against those. Each query then touches at most K·C keys, which scales linearly.
The catch is the picking step. To choose the top K blocks, standard block-sparse attention scores every query against every block summary. That is still O(N²/C), just with a smaller constant. At 128K or 256K tokens, selection dominates. The paper’s own timing shows the baseline BSA selector taking ~313 ms at 256K, versus ~31 ms for PISA. If you’re building a long-context model and you already accepted the engineering pain of sparse attention, the selector is now the bottleneck you actually feel.
The direct comparison point is BSA, a trainable block-sparse baseline built from the selected-attention branch of Native Sparse Attention (NSA) with mean-pooled block summaries. PISA keeps the same overall shape (train sparsely, select K blocks per query, attend to originals) but changes how selection happens.
How It Works
Think of the keys as leaves of a tree. PISA repeatedly mean-pools adjacent blocks to build a pyramid: leaves are the original 64-token key blocks, and each higher level combines a small number (they use branching factor g=2) of children into a coarser summary. The pyramid has O(log N) levels.
Selection walks the tree top-down. At the coarsest level, all candidates are cheap to score. Score them, keep the top K, then descend: only the children of retained blocks become candidates at the next level. Repeat until you reach the leaves, then use those leaf indices for actual sparse attention over the original keys. Because each query only ever sees a bounded candidate set (at most g·K) per level, and there are log N levels, per-query selection is O(log N).
The scoring rule matters. Instead of a dot product against a mean summary, PISA scores a candidate block using LogSumExp over its child summaries. LSE approximates “what would the true attention mass on this block be?” better than mean-pooling does, because attention itself is a softmax (an exponential-weighted sum). They show, via Jensen’s inequality, that LSE-over-children is provably closer to the true raw-key LSE than a mean score is. Two ablations, PISA-1 and PISA-2, truncate LSE at first and second order of a Taylor expansion; the full LSE version wins.
def pisa_select(q, pyramid, K, g): candidates = [root_index] # start at coarsest level for level in range(L, 0, -1): scores = [lse_score(q, pyramid[level][i]) for i in candidates] kept = top_k(candidates, scores, K) if level == 1: return kept # leaf block indices candidates = [child for i in kept for child in children(i)]
To make this fast on a GPU they write Triton kernels. Training/prefill uses a two-stage kernel: stage 1 walks intermediate pyramid levels; stage 2 groups queries that need the same leaf block and scores them together, so each leaf key tile is loaded once and reused across up to Q_tile=4 queries. Decoding uses a single fused kernel instead, because with one query at a time there’s no cross-query key reuse to exploit, and fusion avoids relaunch overhead. An IO-cost analysis in the paper justifies the split.
What They Found
They pretrain decoder-only models at 418M, 1.47B, and 2.67B parameters for 100B tokens at 4K context, then continue for 10B tokens at 16K. Baselines: Full Attention, BSA, Native Sparse Attention (NSA), and HiLS.
•
Language modeling and commonsense reasoning (BoolQ, PIQA, HellaSwag, WinoGrande, ARC, OpenBookQA, SocialIQA): PISA lands within noise of the sparse baselines at all three scales. No method dominates; Full Attention is slightly ahead on averages.
•
Containment/retrieval tasks (SWDE, SQuAD, FDA, TriviaQA, NQ, DROP): PISA has the highest average among sparse methods at all three scales, though Full Attention is still higher. This is the paper’s main quality claim: the coarse-to-fine search doesn’t lose retrieval-relevant blocks.
•
Long-context needle-in-a-haystack on RULER up to 16K, after CPT: PISA is competitive with the best sparse baselines and clearly ahead of BSA on the averaged score.
•
Selection latency: at 64K/128K/256K tokens, PISA achieves 2.86×, 5.31×, and 9.95× speedups over BSA on the selection step alone. Below 16K, BSA is actually faster because its constant factor is small and log N hasn’t paid off yet.
•
Block-selection quality diagnostic: replaying selectors on a frozen Full-Attention model, PISA gets the highest Recall@8 against the true top-K blocks (~91% vs ~86% for BSA) and captures 99.46% of the attention mass the reference set captures. This is the cleanest evidence that LSE scoring picks better blocks, not just faster ones.
Careful reading: the retrieval and RULER wins are modest, and the quality comparison against BSA was done on a frozen Full-Attention checkpoint (not on each method’s own trained model), so it isolates the selector but doesn’t tell you how much of the end-to-end task gap comes from selection quality versus training dynamics.
What’s Useful
•
If you’re training a long-context model from scratch and already committed to sparse attention, PISA is a plausible drop-in for the selection stage. The kernel changes are contained (routing + LSE scoring), and end-task quality is comparable to BSA/NSA at three scales. Worth testing against your own downstream mix.
•
If your context is under ~16K, the paper’s own timings show PISA is slower than BSA at 4K–16K on selection. The log-linear win only shows up at 32K and beyond. Don’t adopt this for short-context serving.
•
If you want to understand why LSE scoring helps, Appendix A is unusually clean: it proves LSE-over-child-means sits between the mean score and the true raw-key LSE, and derives PISA-1/PISA-2 as Taylor truncations. Useful mental model for anyone designing block scorers.
•
The two-stage vs single-stage kernel split is a reusable pattern: reuse key tiles across queries when you have many queries per block (prefill), fuse levels when you have one query per step (decode). The IO analysis in the paper gives you the crossover condition (roughly, when G_Q + C/Q_tile < C).
•
Not established: whether PISA helps as a training-free retrofit onto a dense pretrained model. Every experiment trains with the sparse selector active. The diagnostic on a frozen Full-Attention model is a probe, not a deployment recipe.
No code or checkpoint link is given in the supplied text.
Caveats
•
Model scales top out at 2.67B and pretraining at 100B tokens; the authors flag this as their main limitation. Relative gains over Full Attention can shift a lot at frontier scale.
•
On commonsense benchmarks, PISA is comparable, not better. The contribution is efficiency plus retrieval robustness, not raw quality.
•
The ~10× selection speedup at 256K is for the selection step in isolation, not end-to-end throughput. Attention over the selected blocks and everything else in the forward pass are unchanged.
•
The block-selection quality diagnostic uses a Full-Attention model’s queries and keys as ground truth. That’s a reasonable proxy but not the same as showing PISA picks the blocks its own trained model would benefit most from.
•
Continued-pretraining containment scores at 16K show PISA below BSA at 1.47B and roughly tied at 2.67B, which complicates the “better on retrieval” story from the 4K pretraining table. The picture is genuinely mixed depending on stage and scale.
Topics
Don't miss new content
Log in to follow topics and personalize your feed.
Related topics you might like
Inference Optimization113 episodes
LLM Training134 episodes