SAS: Simple Attention Sparsification via End-to-End Optimization of Context Ranking Review
0. Introduction
Post-training sparse attention은 이미 pretrained된 dense LLM을 크게 다시 학습하지 않고 long-context decode cost를 줄이는 현실적인 경로다. Query마다 중요한 KV block만 고르면 attention computation을 줄일 수 있고, 기존 dense backbone을 유지한 채 selector만 학습할 수 있다.
하지만 learnable sparse attention에는 구조적인 문제가 있다. Selector가 block score를 만든 뒤 hard Top-$K$로 선택하면, selected index는 score에 대해 미분 가능하지 않다. Language modeling loss의 gradient가 selector까지 직접 흐르지 않기 때문에 기존 방법은 dense model의 layer-wise attention distribution을 teacher target으로 distill하는 경우가 많다.
SAS는 여기서 질문을 바꾼다.
Dense model이 많이 본 block과 sparse budget 안에서 최종 prediction에 실제로 유용한 block은 같은가?
Dense attention matching은 original model의 local attention mass를 재현한다. 하지만 fixed sparse budget에서는 layer마다 같은 block을 반복 선택하기보다 서로 보완적인 evidence를 나눠 읽는 편이 더 좋을 수 있고, attention weight가 커도 value contribution이 작을 수 있다.
SAS는 selector의 continuous score를 attention logit 안에 넣어 standard language modeling loss로 context ranking을 end-to-end 최적화한다. Hard Top-$K$는 sparse computation을 위해 유지하지만, selected block에는 soft gate를 남겨 gradient path를 만든다.
한 줄 요약: SAS는 selector score를 log-space gate로 attention softmax 안에 주입해 dense-attention distillation 없이 language modeling loss로 context ranking을 직접 학습하고, tight attention budget에서 reasoning, long-context, agent task 성능을 크게 개선하는 post-training sparse attention method다.
이 논문을 지금 볼 가치가 있는 이유는 다음과 같음.
- Sparse attention의 핵심을 selection이 아니라 fixed budget context ranking 문제로 재정의한다.
- Dense attention distillation이 final prediction utility와 어긋날 수 있다는 명확한 objective mismatch를 지적한다.
- Gate position, activation, continuous ranking, training scope를 controlled ablation으로 분해한다.
- Triton kernel과 SGLang backend를 함께 공개해 training objective를 실제 serving path로 연결한다.
- Math data로 학습한 selector가 long-context and agent benchmark에도 transfer되는 결과를 보여준다.
1. Problem Setting
1-1. Dense attention cost
Current query vector를 $\mathbf{q}$, previous key and value를 $\mathbf{K},\mathbf{V}$라고 하면 dense attention은 다음과 같다.
\[\mathbf{o} = \operatorname{softmax} \left( \mathbf{q}\mathbf{K}^{\top} \right) \mathbf{V}\]Autoregressive decode에서 context length가 계속 늘어나므로 한 token의 attention cost는 $\mathcal{O}(n)$이고, 전체 sequence의 cumulative cost는 $\mathcal{O}(n^2)$가 된다.
Block sparse attention은 context를 block으로 나누고 selected token set $\mathcal{S}$만 읽는다.
\[\mathbf{o} = \operatorname{softmax} \left( \mathbf{q}\mathbf{K}_{\mathcal{S}}^{\top} \right) \mathbf{V}_{\mathcal{S}}\]| Budget가 $ | \mathcal{S} | $로 고정되면 cumulative cost를 $\mathcal{O}(n | \mathcal{S} | )$로 줄일 수 있다. |
1-2. Hard Top-K의 gradient blockage
Selector $\mathcal{R}_{\theta}$가 $C$개 block score $\mathbf{s}$를 만들고 Top-$K$ index를 고른다고 하자.
\[\mathcal{I} = \operatorname{TopK}(\mathbf{s},K)\]Top-$K$ index는 score가 조금 변해도 대부분 그대로다. 따라서 language modeling loss는 어떤 selected block의 attention computation에는 gradient를 주지만, selection order를 만든 selector에는 useful gradient를 직접 주기 어렵다.
기존 learnable sparse attention은 이 문제를 dense attention distillation로 우회한다.
- Original dense attention mass를 block level target으로 만든다.
- Selector가 이 target ranking을 맞추게 한다.
- Inference에서는 Top-$K$ block만 선택한다.
1-3. Dense attention ranking과 sparse utility의 차이
Layer-wise attention matching에는 두 가지 gap이 있다.
1) Cross-layer complementarity를 직접 최적화하지 않는다
각 layer selector가 dense teacher attention을 독립적으로 따라가면 여러 layer가 비슷한 block을 반복 선택할 수 있다. Sparse model 전체의 final loss 관점에서는 layer별로 다른 evidence를 나눠 읽는 것이 더 유용할 수 있다.
2) Value contribution을 보지 않는다
Attention weight가 크다는 사실만으로 final prediction contribution이 크다고 보장할 수 없다. 실제 output은 attention probability와 value vector의 결합이다.
SAS는 selector target을 dense attention weight에서 final language modeling loss로 바꾼다.
2. Core Idea
2-1. Selection 앞의 continuous ranking을 학습한다
SAS는 hard Top-$K$를 없애지 않는다. 대신 Top-$K$가 결정되기 전 selector score가 만드는 continuous ranking을 학습 대상으로 본다.
Historical block score $\mathbf{s}\in\mathbb{R}^{C}$를 positive gate로 변환하고, always-retained current block에는 unit gate를 준다.
\[\mathbf{g} = \phi(\mathbf{s}), \qquad g_0=1\]Historical block gate는 해당 block의 모든 token에 broadcast된다. Gate가 attention computation 안에 들어가면 language modeling loss가 selector score에 gradient를 줄 수 있다.
2-2. Gate를 softmax 안에 log form으로 넣는다
SAS의 main formulation은 다음과 같다.
\[\mathbf{o}_{\mathrm{SAS}} = \operatorname{softmax} \left( \mathbf{q}\mathbf{K}_{\mathcal{S}}^{\top} + \log \mathbf{g}_{\mathcal{S}} \right) \mathbf{V}_{\mathcal{S}}\]Log gate를 attention logit에 더하면 softmax 뒤에서는 multiplicative weight처럼 작동한다.
\[\operatorname{softmax}(a_i+\log g_i) = \frac{g_i e^{a_i}}{\sum_j g_j e^{a_j}}\]따라서 gate는 output을 사후 rescale하는 것이 아니라 attention mass allocation 자체에 참여한다.
2-3. 네 가지 design choice
SAS가 강조하는 핵심은 differentiable path만 만든다고 충분하지 않다는 점이다.
- Gate position
- Gate를 attention softmax 안에 넣는다.
- Gate activation
- Historical score를 softmax로 normalize한다.
- Ranking preservation
- Continuous score difference를 binary mask로 collapse하지 않는다.
- Training scope
- Selected sparse blocks만 사용해도 full-scope training에 가까운 final performance를 얻는다.
3. Architecture / Method
3-1. Overview
| Item | Description |
|---|---|
| Dense backbone | Qwen3-4B, 8B, 14B |
| Selector | SeerAttention-R의 AttnGate architecture |
| Sparse unit | Contiguous context block |
| Block size | 64 tokens |
| Training objective | Standard language modeling loss |
| Gate | Normalized softmax score, log-space inner injection |
| Sparse selection | Query group별 Top-$K$ block |
| Training kernel | FlashAttention-style fused Triton kernel |
| Serving backend | SGLang plus paged KV cache and FlashInfer |
3-2. Gate position
Outer gate는 attention probability가 이미 정해진 뒤 value contribution만 scale한다.
\[\mathbf{o}_{\mathrm{outer}} = \operatorname{softmax} \left( \mathbf{q}\mathbf{K}^{\top} \right) (\mathbf{g}\odot\mathbf{V})\]Inner gate는 normalization에 직접 참여한다.
\[\mathbf{o}_{\mathrm{inner}} = \operatorname{softmax} \left( \mathbf{q}\mathbf{K}^{\top}+\log\mathbf{g} \right) \mathbf{V}\]GPQA-Diamond ablation에서 inner gate가 outer gate보다 크게 앞선다. Selector가 context priority를 학습하려면 gate가 value rescaling이 아니라 attention competition에 들어가야 한다는 evidence다.
3-3. Gate activation
Historical block score를 다음 방식으로 비교한다.
- Softmax-normalized gate
- Independent sigmoid gate
- Raw score injection
Current block은 항상 $g_0=1$이다. Historical gate가 normalize되지 않으면 current block과 historical block의 scale calibration이 불안정해진다. Softmax activation이 sigmoid and raw logit보다 훨씬 안정적인 training을 보인다.
3-4. Continuous ranking preservation
Straight-through estimator로 hard Top-$K$ mask를 쓰면 forward는 discrete하고 backward는 soft score를 흉내 낼 수 있다. 그러나 SAS ablation에서는 continuous gate를 그대로 유지하는 편이 더 안정적이고 final score도 높다.
이 결과는 sparse selector가 selected or not만 배우는 것으로 충분하지 않음을 보여준다. Selected block 안에서도 상대 priority를 학습해야 한다.
3-5. Sparse-scope training
Full-scope training은 모든 context block에 gate를 적용하므로 expensive하다. Sparse-scope는 Top-$K$ selected blocks에만 gate and attention을 계산한다.
초기 convergence는 full scope보다 느리지만 one epoch에서는 comparable result에 도달한다. 실제 post-training cost를 줄이면서 inference condition과 더 가까운 objective를 사용한다는 장점도 있다.
3-6. Fused Triton kernel
Naive implementation은 full attention matrix를 materialize한 뒤 block gate를 더해야 하므로 long sequence에서 memory cost가 크다. SAS kernel은 gate addition을 FlashAttention-style tile computation 안에 fuse한다.
- Full attention matrix를 저장하지 않는다.
- Tile-level $\mathbf{q}\mathbf{K}^{\top}$ 계산 중 log gate를 더한다.
- Backward에서도 selector gradient를 유지한다.
공개 repository는 training recipe뿐 아니라 SGLang-based sparse evaluation backend를 포함한다.
3-7. Inference path
Inference에서는 prefill을 dense attention으로 수행하고, decode에서 block sparse attention을 사용한다.
- Selector가 cached block summary를 scoring한다.
- Query group마다 Top-$K$ block을 고른다.
- Paged KV cache에서 selected block을 읽는다.
- Block-sparse decode kernel을 실행한다.
GQA에서는 query head group이 block selection을 공유해 union access를 줄인다.
4. Training / Data / Recipe
4-1. Training data
Main post-training은 OpenR1-Math-220K의 93.7K examples를 사용한다. Reasoning data로 selector를 학습한 뒤 별도 retraining 없이 다음 task에 평가한다.
- Math reasoning
- GPQA-Diamond
- AIME24 and AIME25
- LongBench-E
- BFCL
- VitaBench
Math-only selector training이 다른 task로 transfer된다는 점은 인상적이지만, domain-independent context ranking을 완전히 증명하는 것은 아니다.
4-2. Training strategy
Main setting은 다음과 같다.
| Item | Setting |
|---|---|
| Backbone sizes | Qwen3-4B, 8B, 14B |
| Backbone update | Frozen |
| Trainable module | AttnGate selector |
| Optimizer | AdamW |
| Learning rate | $10^{-3}$ |
| Schedule | Cosine |
| Global batch size | 32 |
| Maximum sequence length | 32,768 |
| Block size | 64 |
| Training duration | 1 epoch |
| Reported hardware | 8 NVIDIA H20 GPUs |
Dense backbone을 freeze하므로 adaptation target이 selector ranking에 집중된다. Controlled comparison에서는 same backbone, selector architecture, training data를 유지하고 SAS objective와 layer-wise distillation objective를 비교한다.
4-3. Continued pretraining extension
논문은 OLMo3 setting에서 backbone and selector를 함께 학습하는 initial evidence도 제공한다. SAS는 sliding-window attention and HiLS보다 strong average를 보이지만 dense base와 완전히 같지는 않다.
이 결과는 SAS가 post-training trick에만 머물지 않을 가능성을 보여주지만, continued pretraining claim은 Qwen post-training experiment보다 scope가 작다.
4-4. Engineering notes
1) Budget를 token이 아니라 block count와 함께 기록해야 한다
Token budget 2,048는 block size 64에서 32 blocks다. Block size를 바꾸면 same token budget에서도 selector resolution and overhead가 달라진다.
2) Dense prefill and sparse decode를 분리 측정해야 한다
Prompt가 매우 길고 output이 짧으면 dense prefill cost가 지배할 수 있다. Long reasoning처럼 output이 길수록 sparse decode benefit이 커진다.
3) Selector scan cost를 profile해야 한다
Context가 길어질수록 all cached block score와 Top-$K$ operation이 새로운 bottleneck이 된다. Attention FLOPs만 줄여서는 end-to-end speedup을 설명할 수 없다.
4) Gate export and backend version을 pin해야 한다
AttnGate checkpoint, block size, SGLang fork, FlashInfer version이 맞지 않으면 reported path를 재현하기 어렵다.
5) Full-attention fallback이 필요하다
Short context or selector uncertainty가 높은 query에서는 sparse path의 overhead or quality risk가 더 클 수 있다. Dynamic dense-sparse switching을 deployment policy로 둘 수 있다.
5. Evaluation
5-1. Reasoning under tight budgets
SAS의 가장 큰 gain은 low attention budget에서 나타난다.
- 1,024-token budget에서 SeerAttention-R 대비 MATH500은 6.0 to 7.7 points 개선된다.
- 같은 budget에서 GPQA-Diamond는 10.6 to 15.5 points 개선된다.
- Qwen3-4B, 8B, 14B 모두에서 같은 방향을 보인다.
Budget 2,048의 Qwen3-4B example은 다음과 같다.
| Benchmark | SeerAttention-R | SAS |
|---|---|---|
| AIME24 | 55.83 | 68.85 |
| AIME25 | 45.16 | 56.38 |
Budget 4,096에서는 full attention에 가까운 score를 회복한다. Qwen3-4B AIME24에서 SAS는 71.72, full attention은 71.25로 보고된다. Single benchmark의 small overtake를 sparse model이 dense model보다 일반적으로 우수하다는 뜻으로 확대하면 안 된다.
5-2. Long-context understanding
LongBench-E에서 SAS는 tested attention budget 전반에서 SeerAttention-R을 앞선다. 특히 긴 input bucket에서 margin이 커진다.
- Qwen3-14B, budget 2,048, 8K+ bucket: SAS 53.9, SeerAttention-R 51.5
- Qwen3-14B, budget 4,096, overall: SAS 56.2, full attention 56.6
Tight budget에서 ranking quality가 중요하고, budget이 커지면 dense gap이 줄어드는 패턴이다.
5-3. Agentic tasks
BFCL multi-turn function calling에서 다음 result가 보고된다.
- Qwen3-4B, budget 2,048: SAS 32.5, SeerAttention-R 29.0
- Qwen3-14B, budget 4,096: SAS 44.0, full attention 44.5
VitaBench에서도 대부분의 tested setting에서 trainable sparse baseline을 앞서지만, 모든 domain and metric에서 일관된 superiority로 읽기보다 transfer evidence로 보는 편이 적절하다.
5-4. Context ranking analysis
SAS selector는 layer별 dense attention mass recall은 SeerAttention-R보다 낮을 수 있다. 대신 여러 layer가 선택한 block의 union recall은 더 높다.
이 결과는 논문의 objective claim과 잘 맞는다.
- Layer-wise distillation은 각 layer가 original dense attention을 재현하게 한다.
- End-to-end loss는 model 전체가 필요한 evidence를 layer across complementarity로 나눌 수 있게 한다.
즉 local teacher matching이 약해져도 final prediction은 좋아질 수 있다.
5-5. Generation behavior
SAS는 일부 reasoning task에서 average generation length와 truncation rate를 줄인다. Better context ranking이 reasoning loop의 반복을 줄였을 가능성이 있지만, shorter output이 항상 better reasoning을 뜻하지는 않는다.
Accuracy, response length, truncation을 함께 봐야 한다.
5-6. Efficiency
Single-GPU SGLang decode benchmark에서 reported speedup은 다음과 같다.
- Batch 1, 64K context: 2.4x
- Batch 1, 256K context: 4.6x
- Batch 1, 512K context: 5.6x
- Batch 8, 64K context: about 13x
그러나 context가 길어질수록 Top-$K$ selection overhead가 커진다.
- 8K context에서 Top-$K$ share: about 21%
- 512K context에서 Top-$K$ share: about 90%
Sparse attention이 빨라질수록 selector and indexing이 새로운 system bottleneck이 된다는 중요한 결과다.
6. Limitations
- Prefill은 dense다.
- Long prompt and short output workload에서는 end-to-end benefit이 제한될 수 있다.
- Selector는 모든 cached block을 scoring한다.
- Very long context에서 Top-$K$가 latency 대부분을 차지한다.
- Math-only training transfer의 범위가 제한적이다.
- Legal, code repository, multimodal token, retrieval-heavy document task에서 같은 ranking이 유지되는지 확인이 필요하다.
- Block size가 fixed design choice다.
- Fine-grained evidence가 block boundary에 걸리면 coarse selection이 불리할 수 있다.
- Backbone은 Qwen3 중심이다.
- Different attention architecture, MLA, hybrid model, very large MoE에서 재검증이 필요하다.
- Sparse score와 dense score가 근접해도 behavior가 같다는 뜻은 아니다.
- Citation fidelity, rare retrieval, long-range dependency failure를 별도로 평가해야 한다.
- Speedup은 backend and workload dependent다.
- Batch, context, output length, GPU, cache layout이 바뀌면 reported factor도 달라진다.
7. My Take
7-1. Why this matters for my work
SAS의 가장 중요한 메시지는 attention sparsification을 dense attention imitation으로만 풀지 않아도 된다는 점이다. Sparse budget이 실제 deployment constraint라면 selector도 그 budget 안에서 final loss를 줄이는 방향으로 학습해야 한다.
이 관점은 document VLM and long-context RAG에도 직접 연결된다. Layer별 attention map을 teacher로 두기보다, answer token loss, citation loss, grounding loss가 어떤 page block을 남길지 직접 학습하게 만들 수 있다.
7-2. Reuse potential
1) Page-block sparse document attention
Document page or crop를 block으로 두고 answer and grounding loss가 relevant region ranking을 학습하게 할 수 있다.
2) Evidence-aware selector objective
LM loss에 citation entailment or evidence coverage loss를 결합해 selected context가 answer뿐 아니라 provenance도 보존하게 할 수 있다.
3) Hierarchical Top-K
Very long context에서는 coarse retrieval로 candidate block을 줄인 뒤 SAS selector로 final Top-$K$를 고르면 full block scan bottleneck을 줄일 수 있다.
4) Dynamic budget policy
Query difficulty and selector entropy에 따라 1,024, 2,048, 4,096 budget를 바꾸는 adaptive compute를 설계할 수 있다.
5) Cross-layer diversity regularization
SAS가 암묵적으로 얻은 cross-layer complementarity를 explicit diversity objective or coverage metric으로 측정할 수 있다.
7-3. Follow-up papers
- SeerAttention
- SeerAttention-R
- Native Sparse Attention
- Quest
- RetrievalAttention
- MoBA
- MInference
- DuoAttention
8. Summary
- SAS는 sparse attention selector를 dense attention distribution이 아니라 language modeling loss로 직접 학습한다.
- Continuous selector score를 log-space gate로 attention softmax 안에 넣어 gradient path를 만든다.
- Normalized gate, continuous ranking, sparse-scope training, fused Triton kernel이 핵심 design choice다.
- Tight budget에서 reasoning, long-context, agent benchmark의 gain이 가장 크고 4,096-token budget에서는 dense performance에 근접한다.
- Dense prefill과 all-block Top-$K$ scan이 남아 있어 다음 bottleneck은 selector and indexing system이다.
댓글남기기