Sparser Block-Sparse Attention via Token Permutation
Xinghao Wang, Pengyu Wang, Dong Zhang, Chenkun Tan, Shaojun Zhou, Zhaoxiang Liu, Shiguo Lian, Fangxu Liu, Kai Song, Xipeng Qiu
TL;DR
This work tackles the $O(N^2)$ cost of self-attention for long-context LLMs by introducing Permuted Block-Sparse Attention (PBS-Attn), a plug-and-play method that rearranges tokens to produce a more favorable block-sparse structure while preserving causality through segmented permutation. It leverages attention permutation invariances to implement a query-aware key permutation that clusters salient key tokens within segments, enabling sparser block interactions without sacrificing accuracy. The authors provide a formal treatment of permutation properties, a practical segmentation strategy, and a custom permuted-FlashAttention kernel, along with extensive experiments on LongBench and LongBenchv2 showing near-full-attention performance with up to $2.75×$ end-to-end speedups for long-context prefilling. This approach offers a scalable and hardware-friendly path to more efficient long-context LLMs with broad practical impact for real-world long-sequence tasks.
Abstract
Scaling the context length of large language models (LLMs) offers significant benefits but is computationally expensive. This expense stems primarily from the self-attention mechanism, whose $O(N^2)$ complexity with respect to sequence length presents a major bottleneck for both memory and latency. Fortunately, the attention matrix is often sparse, particularly for long sequences, suggesting an opportunity for optimization. Block-sparse attention has emerged as a promising solution that partitions sequences into blocks and skips computation for a subset of these blocks. However, the effectiveness of this method is highly dependent on the underlying attention patterns, which can lead to sub-optimal block-level sparsity. For instance, important key tokens for queries within a single block may be scattered across numerous other blocks, leading to computational redundancy. In this work, we propose Permuted Block-Sparse Attention (\textbf{PBS-Attn}), a plug-and-play method that leverages the permutation properties of attention to increase block-level sparsity and enhance the computational efficiency of LLM prefilling. We conduct comprehensive experiments on challenging real-world long-context datasets, demonstrating that PBS-Attn consistently outperforms existing block-sparse attention methods in model accuracy and closely matches the full attention baseline. Powered by our custom permuted-FlashAttention kernels, PBS-Attn achieves an end-to-end speedup of up to $2.75\times$ in long-context prefilling, confirming its practical viability. Code available at https://github.com/xinghaow99/pbs-attn
