Command Palette
Search for a command to run...
対数線形計算量を実現するブロックスパース注意機構
対数線形計算量を実現するブロックスパース注意機構
Bohao Tang Zhen Qin Yuqi Pan Zheng Li Pengfei Liu
概要
言語モデルを長文脈へ拡張することは、自己注意の二次計算量によって制限される。ブロックスパース注意は効率的な代替手段となるが、保持するブロックの選択が依然としてボトルネックとなる。従来のブロック選択では、すべてのクエリ・ブロック対をスコアリングする必要があり、そのため系列長に対して二次の計算量にとどまる。この問題に対処するため、我々はピラミッド型Top-K選択戦略を用いるブロックスパース注意機構PISAを提案する。主なアイデアは、異なるレベルにわたって候補を段階的に絞り込み、最も関連性の高いキーをより効率的に見つけることである。具体的には、キーの粗密階層を構築し、最も粗いレベルから選択を行う。各レベルでは、有界な候補集合に対してLogSumExpスコアリングを適用し、次のより細かいレベルの候補を選択し、最も細かいレベルに到達するまで続ける。プーリングによりO(log N)レベルのキーを構築し、全体の計算量はO(N log N)となる。ここでNは系列長を表す。我々は学習と推論の両方のためのハードウェアを考慮したTritonカーネルを開発し、階層的ルーティングとLogSumExpスコアリングを融合し、クエリ・キースコア行列を明示的に生成しない。さらに、言語モデリングタスクで本手法を評価する。ベースラインと比較して、本手法は常識推論などのベンチマークで同等の性能を達成しつつ、検索タスクではより優れた結果を示す。
One-sentence Summary
Researchers at Shanghai Jiao Tong University, Shanghai Innovation Institute, and ByteDance Seed propose PISA, a block-sparse attention mechanism that uses pyramid Top-K selection over a coarse-to-fine key hierarchy with bounded LogSumExp scoring to achieve O(N log N) complexity, and they develop hardware-aware Triton kernels for training and inference, yielding performance comparable to baselines on language modeling benchmarks and better results on retrieval tasks.
Key Contributions
- The paper introduces PISA, a block-sparse attention mechanism that uses pyramid Top-K selection over a coarse-to-fine hierarchy of keys and LogSumExp scoring at each level, reducing block selection complexity to O(N log N).
- Hardware-aware Triton kernels for training and inference fuse hierarchical routing and LogSumExp scoring, maintain a bounded candidate set, and avoid materializing the dense query-key score matrix.
- Experiments on language modeling show that PISA achieves comparable performance to conventional block-sparse attention baselines on commonsense reasoning, delivers better results on retrieval tasks, and attains lower selection latency than single-level block selection for sequences up to 256K tokens.
Introduction
Scaling language models to longer contexts is important, but full self-attention grows quadratically with sequence length. Block-sparse attention reduces this cost by restricting each query to selected key blocks; however, conventional Top-K block selection still scans every candidate key block, creating an O(N2/C) selection bottleneck for long sequences. To address this, the authors propose PISA, a pyramid sparse attention method that builds a multilevel pooled key hierarchy and performs coarse-to-fine Top-K selection with LogSumExp scoring. This reduces block-selection complexity to O(NlogN) and uses hardware-aware Triton kernels that fuse hierarchical routing while avoiding dense query-key score matrices.
Method
The authors introduce PISA as a pyramid block sparse attention mechanism. It combines three main components: a block sparse attention backbone, a hierarchical top-K block selector built over a coarse-to-fine key pyramid, and dedicated training, prefill, and decoding kernels.
Block Sparse Attention
Keys and values are partitioned into contiguous blocks. For a query qt, the block selector retains only a subset of key blocks, and attention is computed over the corresponding original keys and values:
ot=∑j∈Ttexp(qt⊤kj/d)∑j∈Ttexp(qt⊤kj/d)vj.A flat block sparse attention method scores M=⌈N/C⌉ block summaries for every query, which still costs O(N2d/C) for full-sequence scoring. PISA instead organizes block selection as a hierarchical search.
Pyramid Key-Block Hierarchy
The key hierarchy is constructed from fine to coarse. Level 0 contains the original keys as singleton blocks:
M0=N,Ki(0)=[ki].The first merging step uses g0=C and produces leaf blocks. For all higher levels, the merging factor is gℓ=g=2:
Mℓ+1=⌈gℓMℓ⌉, Ki(ℓ+1)=Concatr∈Chℓ(i)Kr(ℓ).Each block Ki(ℓ) has a summary vector kˉi(ℓ). Starting from kˉi(0)=ki, the summaries are computed recursively by mean pooling:
kˉi(ℓ+1)=∣Chℓ(i)∣1r∈Chℓ(i)∑kˉr(ℓ).Because the hierarchy uses a fixed merging factor g=2 after the leaf level, it contains L=O(logN) levels. Queries remain token-wise throughout the selection process.
Coarse-to-Fine Top-K Selection
Selection proceeds from the coarsest level L down to the leaf level 1. For each query, the candidate set at the coarsest level is initialized as:
At(L)={1}.At level ℓ, PISA scores only the candidate blocks in At(ℓ), selects the top K blocks, and expands the retained blocks into their children:
At(ℓ−1)=i∈It(ℓ)⋃Chℓ−1(i).Since each retained block has at most g children for ℓ>1, the candidate set size is bounded by gK at the next level. If a level contains at most K valid candidates, all candidates are retained without scoring. The final leaf-level selection set Tt=Tt(1) determines the original key and value blocks used by the sparse attention layer.
For scoring, PISA uses a LogSumExp block score. The score of candidate block Ki(ℓ) for query qt is:
st,i(ℓ)=logr∈Chℓ−1(i)∑exp(dqt⊤kˉr(ℓ−1)),ℓ≥1.At the leaf level, the children are original keys, so a leaf score is the exact raw-key LogSumExp score of that block. At intermediate levels, each candidate is scored from its child summaries, and each score reduces at most g logits. This keeps the scoring cost small while making leaf-level selection consistent with the attention log-sum-exp computation.
Kernel Implementation
For training and prefill, PISA uses a two-stage kernel design.
In the first stage, the kernel processes intermediate levels from L to 2. It parallelizes over query positions and key/value heads. Query heads within the same GQA group share the same selected key blocks. For a key/value head h, the kernel computes LSE scores for all query heads in H(h) and sums them:
uh,t,i(ℓ)=h′∈H(h)∑sh′,t,i(ℓ).It then applies top-K selection on the summed scores and expands the retained blocks into their children to form the next candidate set. The candidate indices remain local to the program throughout this loop.
After the first stage, each query has at most gK candidate leaf blocks. The second stage groups queries that need to score the same candidate key block into clusters. Each cluster is divided into tiles of at most Qtile queries, with Qtile=4 in the implementation. Each kernel program loads a key tile of shape d×C once and reuses it across a query tile of shape (QtileGQ)×d, where GQ=∣H(h)∣. It computes
(QtileGQ)×d×d×C⟶(QtileGQ)×C,applies LSE over keys for each query head, and then sums over the GQ query heads. A final kernel independently applies top-K selection for each query and key/value head.
During decoding, PISA caches the mean pyramid and updates only the current leaf mean and the affected ancestor path for each new token. Block selection uses a single fused kernel that runs from the coarsest level down to ℓ=1, directly scoring the original key blocks. This avoids the overhead of launching multiple kernels per decoding step.
Complexity
Because the pyramid hierarchy has O(logN) levels and each query scores at most gK candidate blocks per level, the selection cost per query is O(logN). Full-sequence training or prefill therefore costs O(NlogN), while decoding costs O(logN) per step on average.
The two-stage training/prefill kernel is chosen mainly for input/output efficiency. By grouping queries that share a candidate key block, the second stage reuses each loaded key block across multiple queries. For the implementation settings GQ=16, Qtile=4, and C=64,
GQ+QtileC=32<64=C,so the two-stage design reduces leaf-level query/key IO compared with a single-stage design. During decoding, clusters contain only one query, so there is no cross-query key reuse. The fused single-stage kernel reuses the query vectors already loaded during intermediate selection and avoids separate grouping and final-selection launches.
Experiment
The experiments compare PISA with Full Attention, BSA, NSA, and HiLS at 418M, 1.47B, and 2.67B scales under matched training, with 100B tokens at 4K sequence length followed by 10B tokens of continued pretraining at 16K. PISA performs comparably to sparse baselines on language modeling and commonsense reasoning, achieves the highest average containment accuracy among sparse methods at all scales though Full Attention remains higher, and shows strong long-context retrieval on RULER. Ablations indicate that the LSE-based scoring improves Top-K block selection quality and downstream performance relative to PISA-1, PISA-2, and BSA, while block-selection efficiency becomes advantageous over BSA at very long contexts.
The table compares training settings and computational complexity for sparse attention methods. HiP is notable as a training-free method with subquadratic prefill and logarithmic decode complexity, while HISA and most trainable methods use quadratic prefill and linear decode complexity. LLSA is the trainable exception with subquadratic prefill, but its cost is reported per diffusion attention pass rather than autoregressive prefill. HiP achieves lower prefill and decode complexity than HISA among the listed training-free methods, with subquadratic prefill and logarithmic decode. Most trainable methods share quadratic prefill and linear decode complexity; LLSA has subquadratic prefill but is reported for diffusion attention passes rather than autoregressive prefill.
PISA performs comparably to BSA, NSA, and HiLS on language modeling and commonsense reasoning after 100B tokens of pretraining. Among sparse methods, PISA achieves the highest average accuracy on containment tasks across evaluated scales, while full attention retains a higher containment average. PISA also shows lower training loss and higher average containment accuracy than PISA-1, PISA-2, and BSA. PISA reaches the best containment average among the compared sparse methods across model scales. Full attention still leads containment accuracy, while PISA matches sparse baselines on language modeling and commonsense reasoning. PISA improves over PISA-1, PISA-2, and BSA in training loss and average containment accuracy.
After continued pretraining at 16K, full attention has the highest average RULER needle-in-a-haystack accuracy for the 2.67B models. Among sparse methods, PISA and PISA-2 lead overall and are nearly tied, with PISA-2 slightly ahead on average and both above NSA and BSA. Accuracy declines with context length across methods, especially on multi-key, multi-query, and multi-value tasks. Full attention records the top average accuracy, while PISA-2 and PISA are the strongest sparse methods and nearly tied. PISA variants outperform full attention on several short-context multi-key and multi-query settings, but all sparse methods drop more sharply by 16K, especially on multi-key and multi-value tasks.
The experiments compare sparse attention methods on efficiency, language modeling and commonsense reasoning after 100B-token pretraining, and long-context retrieval after continued pretraining at 16K. PISA performs comparably to BSA, NSA, and HiLS on language modeling and commonsense reasoning, and it achieves the highest average containment accuracy among sparse methods, though full attention remains stronger overall. In long-context RULER evaluation, PISA and PISA-2 are the strongest sparse methods and nearly tied, with some short-context advantages over full attention, but all sparse methods degrade more sharply by 16K. Complexity results also show that training-free HiP offers subquadratic prefill and logarithmic decode, while most trainable methods use quadratic prefill and linear decode.