long-context language-modeling commonsense-reasoning triton-kernels retrieval-tasks hierarchical-routing block-sparse-attention log-linear-complexity
Abstract
Scaling language models to long contexts is limited by the quadratic cost of self-attention. Block sparse attention offers an efficient alternative, but selecting the retained blocks remains a bottleneck. Conventional block selection requires scoring all query-block pairs and therefore remains quadratic in sequence length. To address this issue, we propose PISA, a block-sparse attention mechanism that employs a pyramid Top-K selection strategy. The main idea is to gradually narrow down the candidates across different levels, making it more efficient to find the most relevant keys. Specifically, we construct a coarse-to-fine hierarchy of keys and perform selection from the coarsest level. At each level, LogSumExp scoring is applied to a bounded candidate set to select candidates for the next finer level, continuing until the finest level is reached. Through pooling, we construct O(log N) levels of keys, yielding an overall complexity of O(Nlog N), where N denotes the sequence length. We develop hardware-aware Triton kernels for both training and inference, fusing hierarchical routing and LogSumExp scoring without materializing the query-key score matrix. We further evaluate our method on language modeling tasks. Compared with the baseline, our method achieves comparable performance on benchmarks such as commonsense reasoning while delivering better results on retrieval tasks.
한국어 요약
한 줄 요약
PISA는 로그-선형 복잡도를 갖는 블록 스파스 어텐션 메커니즘으로, 긴 문맥 처리 성능을 향상시킨다.
핵심 기여도
- PISA는 블록 스파스 어텐션에서 블록 선택 과정의 제곱 복잡도 문제를 해결한다.
- 피라미드 형태의 Top-K 선택 전략을 도입하여 O(N log N)의 복잡도를 달성한다.
- 하드웨어 최적화된 Triton 커널을 개발하여 훈련 및 추론 과정에서 효율성을 높인다.
- 기존 베이스라인 대비 검색 작업 성능을 개선한다.
핵심 아이디어
PISA는 기존 블록 스파스 어텐션에서 블록 선택 시 발생하는 제곱 복잡도 문제를 해결하기 위해, 코어 아이디어로 **피라미드 형태의 Top-K 선택 전략**을 제안한다. 이는 어텐션 키를 **가장 거친 수준부터 점진적으로 세부 수준으로 이동하며 후보군을 줄이는 방식**이다. 각 수준에서 **LogSumExp 점수 계산**을 통해 다음 단계의 후보를 선정함으로써, 전체적으로 O(log N) 개의 키 레벨을 구성하고, 최종적으로 O(N log N)의 복잡도를 달성한다. 이는 기존의 모든 쿼리-블록 쌍을 점수화하는 방식에 비해 효율적이다.
기술적 접근법
- **PISA**는 **피라미드 Top-K 선택 전략**을 사용하여 블록 스파스 어텐션의 효율성을 높인다.
- **LogSumExp 점수 계산**을 통해 각 레벨에서 후보 키를 선정하고, **O(log N)** 개의 레벨을 구성한다.
- **Triton 커널**을 사용하여 **계층적 라우팅**과 **LogSumExp 점수 계산**을 **쿼리-키 점수 행렬 생성 없이** 수행한다.
- **하드웨어 최적화**를 통해 훈련 및 추론 과정에서의 성능을 향상시킨다.
주요 결과
- **PISA**는 **commonsense reasoning** 벤치마크에서 기존 베이스라인과 **비교 가능한 성능**을 보인다.
- **검색 작업**(retrieval tasks)에서는 베이스라인 대비 **더 우수한 결과**를 달성한다.
- **복잡도**는 O(N log N)으로, 기존의 O(N²)에 비해 **가속화**를 달성한다.
의의 및 한계
PISA는 긴 문맥 처리를 위한 어텐션 메커니즘의 효율성을 높이는 데 기여하며, 특히 **검색 작업 성능 향상**이 실용적 가치를 제공한다. 또한, **하드웨어 최적화**를 통해 실제 시스템에 적용 가능성을 높인다. 그러나, **모든 어텐션 작업에서의 성능 개선 여부**는 명시되지 않았으며, **다양한 어텐션 패턴에 대한 일반화 가능성**도 추가 연구가 필요하다.
실용적 활용
PISA는 **긴 문맥을 처리하는 대형 언어 모델**에서 유용하게 활용될 수 있으며, 특히 **검색 엔진**, **문서 이해 시스템**, **대화형 AI** 등에서 성능 향상 효과를 기대할 수 있다.