PISA cuts block-sparse attention selection to O(N log N) with pyramid Top-K, running 9.95x faster than BSA at 256K
Synopsis
Researchers from Shanghai Jiao Tong University and ByteDance Seed propose PISA, a block-sparse attention mechanism that uses coarse-to-fine pyramid Top-K selection with LogSumExp scoring and hardware-aware Triton kernels to reduce block-selection complexity from O(N²/C) to O(N log N); across 418M, 1.47B, and 2.67B scales it matches BSA, NSA, and HiLS on language modeling and commonsense reasoning, achieves the highest average accuracy among sparse methods on six containment tasks, and speeds up block selection over BSA by 2.86x, 5.31x, and 9.95x at 64K, 128K, and 256K.
Interpretation
PISA replaces 'score every key block for every query' in block-sparse attention Top-K selection with coarse-to-fine narrowing over a pyramid of key blocks. Conventional block-sparse attention such as BSA must score all N/C key blocks per query, leaving the selection stage at O(N²/C) overall; PISA builds O(log N) levels of key representations via mean pooling, scores at most gK candidates per level with LogSumExp, keeps the Top-K, and expands them to the next finer level, giving O(log N) routing per query and O(N log N) over the sequence. The paper provides the full complexity derivation: L=O(log N) pyramid levels, bounded candidates per level, plus an IO-cost comparison of the training/prefill and decoding kernel designs (with GQ=16, Qtile=4, C=64 giving GQ+C/Qtile=32<64=C).
The authors implement hardware-aware Triton kernels for both training/prefill and decoding that fuse hierarchical routing and LogSumExp scoring without materializing the query-key score matrix. Training/prefill uses a two-stage kernel: stage one parallelizes over query positions and KV heads, summing scores across query heads sharing a KV head; stage two clusters queries needing the same candidate key block and tiles them with Qtile=4 to reuse each loaded key block. Decoding uses a single-stage kernel that caches the mean pyramid and updates only the new token's leaf mean and ancestor path. The paper gives an analytic comparison of leaf-level Q/K IO cost for the two-stage versus single-stage designs, showing the two-stage design has lower IO under GQ=16, Qtile=4, C=64, while decoding has one query per cluster and no cross-query reuse, favoring the single-stage kernel.
At 418M, 1.47B, and 2.67B scales, PISA matches BSA, NSA, and HiLS on language modeling and commonsense reasoning while doing better on retrieval-style tasks. After 100B tokens of pretraining at 4K, PISA achieves the highest average accuracy among sparse methods on six containment tasks (SWDE, SQuAD, FDA, TQA, NQ, DROP): 41.77 at 418M, 49.83 at 1.47B, and 52.14 at 2.67B, though Full Attention remains higher; on RULER retrieval after 16K continued pretraining, the 2.67B PISA model averages 62.80 versus BSA's 54.99 and NSA's 61.07, below Full Attention's 66.24. All three scales use the same decoder-only backbone and matched training settings, with sparse methods at C=64 and K=8 (K=32 for continued pretraining), reporting training loss, perplexity, multiple-choice accuracy, and containment accuracy; after continued pretraining PISA's average containment accuracy is lower than BSA's.
Block-selection efficiency improves markedly with sequence length: PISA's Q4 implementation is 2.86x, 5.31x, and 9.95x faster than BSA at 64K, 128K, and 256K. BSA is faster from 4K to 16K, but Q4 has the lowest latency from 32K onward; Q4 also speeds up over the per-query implementation by 1.35x, 1.30x, and 1.33x at 64K, 128K, and 256K. The latency benchmark uses random BF16 queries and keys with a fixed seed, batch size 1, 32 query heads, two KV heads, d=64, C=64, K=8, and g=2; timings include summary construction and all selection stages but exclude sparse attention itself.
Perspective
This work targets training and inference settings that need long-context language modeling and retrieval, for decoder-only backbones with GQA grouping and fixed block size and budget; it delivers complexity and latency improvements in the selection stage plus comparable performance to sparse baselines under those settings. For teams aiming to cut long-sequence prefill and decoding selection overhead, the paper offers a reusable pyramid-selection idea, a LogSumExp scoring form, and Triton kernel designs (two-stage for training/prefill, single-stage for decoding), along with a block-selection quality diagnostic (Recall@K, captured mass, attention mass ratio) for later comparisons.
The paper states that computational constraints limited the model sizes and pretraining data budgets explored, and notes both can substantially affect performance and relative gains, so behavior at larger scale and longer training remains an open question. After continued pretraining, PISA's average containment accuracy is lower than BSA's, while Full Attention still leads on multiple-choice and containment averages, indicating the trade-offs of sparsification differ across stages. The block-selection diagnostic reports PISA's Recall@8 at 90.95 and attention mass ratio at 99.46, but it is run on fixed Full-Attention weights and the reference set aggregates scores differently from PISA, so the two can rank blocks differently. In addition, this reading is of the full paper text with figures and some tables rendered as text; verifying exact curve values still requires the original figures.
