HyperAIHyperAI

Command Palette

Search for a command to run...

الانتباه المتناثر القائم على الكتل بتعقيد لوغاريتمي خطي

Bohao Tang Zhen Qin Yuqi Pan Zheng Li Pengfei Liu

الملخص

تشكل التكلفة التربيعية للانتباه الذاتي قيدًا أمام توسيع نماذج اللغة إلى سياقات طويلة. يوفر الانتباه المتناثر القائم على الكتل بديلاً فعالاً، لكن اختيار الكتل المُبقاة يظل عنق زجاجة. يتطلب اختيار الكتل التقليدي تقييم درجات جميع أزواج الاستعلام–الكتلة، ومن ثم يبقى تربيعيًا في طول التسلسل. لمعالجة هذه المسألة، نقترح PISA، وهي آلية انتباه متناثرة قائمة على الكتل تستخدم استراتيجية اختيار هرمية من نوع Top-K. الفكرة الرئيسية هي تضييق نطاق المرشحين تدريجيًا عبر مستويات مختلفة، مما يجعل إيجاد المفاتيح الأكثر صلة أكثر كفاءة. وتحديدًا، نبني تسلسلًا هرميًا للمفاتيح من الخشن إلى الدقيق، ونجري الاختيار بدءًا من المستوى الأكثر خشونة. في كل مستوى، يُطبَّق تسجيل 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)O(N^2/C)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(Nlog⁡N)O(N \log N)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-KKK 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 qtq_tqt​, 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)vj∑j∈Ttexp⁡(qt⊤kj/d).o_t = \frac{\sum_{j \in \mathcal{T}_t} \exp\left(q_t^\top k_j / \sqrt{d}\right) v_j} {\sum_{j \in \mathcal{T}_t} \exp\left(q_t^\top k_j / \sqrt{d}\right)}.ot​=∑j∈Tt​​exp(qt⊤​kj​/d​)∑j∈Tt​​exp(qt⊤​kj​/d​)vj​​.

A flat block sparse attention method scores M=⌈N/C⌉M = \lceil N/C \rceilM=⌈N/C⌉ block summaries for every query, which still costs O(N2d/C)O(N^2 d / C)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 000 contains the original keys as singleton blocks:

M0=N,Ki(0)=[ki].M_0 = N, \qquad K_i^{(0)} = [k_i].M0​=N,Ki(0)​=[ki​].

The first merging step uses g0=Cg_0 = Cg0​=C and produces leaf blocks. For all higher levels, the merging factor is gℓ=g=2g_\ell = g = 2gℓ​=g=2:

Mℓ+1=⌈Mℓgℓ⌉,M_{\ell+1} = \left\lceil \frac{M_\ell}{g_\ell} \right\rceil,Mℓ+1​=⌈gℓ​Mℓ​​⌉, Ki(ℓ+1)=Concatr∈Chℓ(i)Kr(ℓ).K_i^{(\ell+1)} = \mathrm{Concat}_{r \in \mathrm{Ch}_\ell(i)} K_r^{(\ell)}.Ki(ℓ+1)​=Concatr∈Chℓ​(i)​Kr(ℓ)​.

Each block Ki(ℓ)K_i^{(\ell)}Ki(ℓ)​ has a summary vector kˉi(ℓ)\bar{k}_i^{(\ell)}kˉi(ℓ)​. Starting from kˉi(0)=ki\bar{k}_i^{(0)} = k_ikˉi(0)​=ki​, the summaries are computed recursively by mean pooling:

kˉi(ℓ+1)=1∣Chℓ(i)∣∑r∈Chℓ(i)kˉr(ℓ).\bar{k}_i^{(\ell+1)} = \frac{1}{|\mathrm{Ch}_\ell(i)|} \sum_{r \in \mathrm{Ch}_\ell(i)} \bar{k}_r^{(\ell)}.kˉi(ℓ+1)​=∣Chℓ​(i)∣1​r∈Chℓ​(i)∑​kˉr(ℓ)​.

Because the hierarchy uses a fixed merging factor g=2g=2g=2 after the leaf level, it contains L=O(log⁡N)L = O(\log N)L=O(logN) levels. Queries remain token-wise throughout the selection process.

Coarse-to-Fine Top-KKK Selection

Selection proceeds from the coarsest level LLL down to the leaf level 111. For each query, the candidate set at the coarsest level is initialized as:

At(L)={1}.\mathcal{A}_t^{(L)} = \{1\}.At(L)​={1}.

At level ℓ\ellℓ, PISA scores only the candidate blocks in At(ℓ)\mathcal{A}_t^{(\ell)}At(ℓ)​, selects the top KKK blocks, and expands the retained blocks into their children:

At(ℓ−1)=⋃i∈It(ℓ)Chℓ−1(i).\mathcal{A}_t^{(\ell-1)} = \bigcup_{i \in \mathcal{I}_t^{(\ell)}} \mathrm{Ch}_{\ell-1}(i).At(ℓ−1)​=i∈It(ℓ)​⋃​Chℓ−1​(i).

Since each retained block has at most ggg children for ℓ>1\ell > 1ℓ>1, the candidate set size is bounded by gKgKgK at the next level. If a level contains at most KKK valid candidates, all candidates are retained without scoring. The final leaf-level selection set Tt=Tt(1)\mathcal{T}_t = \mathcal{T}_t^{(1)}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(ℓ)K_i^{(\ell)}Ki(ℓ)​ for query qtq_tqt​ is:

st,i(ℓ)=log⁡∑r∈Chℓ−1(i)exp⁡(qt⊤kˉr(ℓ−1)d),ℓ≥1.s_{t,i}^{(\ell)} = \log \sum_{r \in \mathrm{Ch}_{\ell-1}(i)} \exp\left(\frac{q_t^\top \bar{k}_r^{(\ell-1)}}{\sqrt{d}}\right), \quad \ell \geq 1.st,i(ℓ)​=logr∈Chℓ−1​(i)∑​exp(d​qt⊤​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 ggg 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 LLL to 222. 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 hhh, the kernel computes LSE scores for all query heads in H(h)\mathcal{H}(h)H(h) and sums them:

uh,t,i(ℓ)=∑h′∈H(h)sh′,t,i(ℓ).u_{h,t,i}^{(\ell)} = \sum_{h' \in \mathcal{H}(h)} s_{h',t,i}^{(\ell)}.uh,t,i(ℓ)​=h′∈H(h)∑​sh′,t,i(ℓ)​.

It then applies top-KKK 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 gKgKgK 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 QtileQ_{\text{tile}}Qtile​ queries, with Qtile=4Q_{\text{tile}}=4Qtile​=4 in the implementation. Each kernel program loads a key tile of shape d×Cd \times Cd×C once and reuses it across a query tile of shape (QtileGQ)×d(Q_{\text{tile}} G_Q) \times d(Qtile​GQ​)×d, where GQ=∣H(h)∣G_Q = |\mathcal{H}(h)|GQ​=∣H(h)∣. It computes

(QtileGQ)×d×d×C⟶(QtileGQ)×C,(Q_{\text{tile}} G_Q) \times d \times d \times C \longrightarrow (Q_{\text{tile}} G_Q) \times C,(Qtile​GQ​)×d×d×C⟶(Qtile​GQ​)×C,

applies LSE over keys for each query head, and then sums over the GQG_QGQ​ query heads. A final kernel independently applies top-KKK 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\ell=1ℓ=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(log⁡N)O(\log N)O(logN) levels and each query scores at most gKgKgK candidate blocks per level, the selection cost per query is O(log⁡N)O(\log N)O(logN). Full-sequence training or prefill therefore costs O(Nlog⁡N)O(N \log N)O(NlogN), while decoding costs O(log⁡N)O(\log N)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=16G_Q = 16GQ​=16, Qtile=4Q_{\text{tile}} = 4Qtile​=4, and C=64C = 64C=64,

GQ+CQtile=32<64=C,G_Q + \frac{C}{Q_{\text{tile}}} = 32 < 64 = C,GQ​+Qtile​C​=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.


بناء الذكاء الاصطناعي بالذكاء الاصطناعي

من الفكرة إلى الإطلاق — سرّع تطوير الذكاء الاصطناعي الخاص بك مع المساعدة البرمجية المجانية بالذكاء الاصطناعي، وبيئة جاهزة للاستخدام، وأفضل أسعار لوحدات معالجة الرسومات.

البرمجة التعاونية باستخدام الذكاء الاصطناعي
وحدات GPU جاهزة للعمل
أفضل الأسعار

HyperAI Newsletters

اشترك في آخر تحديثاتنا
سنرسل لك أحدث التحديثات الأسبوعية إلى بريدك الإلكتروني في الساعة التاسعة من صباح كل يوم اثنين
مدعوم بواسطة MailChimp