SAS는 언어모델 손실로 문맥 순위를 직접 학습하는 사후 어텐션 희소화다.

SAS: Simple Attention Sparsification via End-to-End Optimization of Context Ranking

HF Daily2609.13141

Zhiwei Li, Lei Zhu, Hao Gu2026-09-11조회 4

무엇인가

긴 문맥 추론은 LLM의 핵심 병목이다. 자기회귀 생성은 새 토큰마다 이전 문맥 전체에 어텐션해야 하고, 누적 비용은 문맥 길이의 제곱으로 늘어난다. 이미 dense 어텐션으로 배포된 모델이 많기 때문에, 아키텍처 변경이나 재학습 없이 사후(post-training)에 어텐션을 희소화하는 접근이 선호된다. 문제의 핵심은 제한된 어텐션 예산 아래에서 각 쿼리에 가장 유용한 문맥 단위(토큰 또는 블록)를 고르는 것이다. 기존 학습형 방법들은 경량 선택기(selector)로 점수를 매긴 뒤 하드 Top-K로 블록을 고르는데, 이 선택 단계가 미분 불가능해 언어모델 손실의 그래디언트가 선택기로 흐르지 못한다. 그래서 대개 원본 dense 모델의 층별 어텐션 분포를 증류한다. 그러나 이 지도 신호는 "원본 모델이 어디를 보는가"를 순위로 가르칠 뿐, 고정 예산에서 "최종 예측에 얼마나 기여하는가"와 직접 정렬되지 않는다. 저자들은 이 불일치를 순위 불일치(ranking misalignment)라 부르고, 층별 목표가 층 간 상호보완성을 놓치며 어텐션 가중치만 맞출 뿐 V 행렬이 예측에 미치는 영향을 무시한다는 두 한계를 지적한다.

어떻게 동작하나

제안 방법 SAS(Simple Attention Sparsification)는 선택기 점수를 학습 중 어텐션 계산에 직접 참여시켜, 표준 언어모델 손실로 선택기를 end-to-end 학습한다. 현재 블록 B0는 항상 유지하고, 나머지 C개 과거 블록 H에 대해 선택기가 점수 s를 낸다. 이를 양수 게이트 g = φ(s), g0 = 1로 변환하고, 게이트를 블록 내 모든 토큰에 브로드캐스트한 뒤 어텐션 로짓에 로그 형태로 더한다. 학습 시 어텐션은 o_SAS = softmax(qK_S^T + log g_S) V_S로 계산되며, 선택된 블록 집합 S = B0 ∪ (Top-K 블록들)이다. 추론 시에는 학습된 연속 순위를 이산화해 블록을 고른다. 즉 하드 Top-K를 연속적인 로그 공간 게이트로 완화한 것이 이 방법의 전부이며, 교사 어텐션도 보조 증류 손실도 필요 없다.

무엇과 다른가

저자들은 이 단순한 설계가 실제로 작동하려면 네 가지 선택이 중요하다고 밝힌다. 첫째, 게이트는 softmax 안쪽에 로그 형태로 넣어야 한다. softmax 바깥에 두면 어텐션 확률이 이미 고정된 채 V만 재조정되지만, 안쪽에 두면 블록 간 어텐션 질량 재분배를 직접 학습한다. 논문은 게이트 그래디언트가 바깥쪽에서는 dg = Σ p_i dO^T v_i인 반면 안쪽에서는 dg = Σ (p̃_i/g_m) dO^T (v_i − o)로, 후자가 v_i − o를 통해 상대적 중요도 신호를 준다는 점을 보인다. 둘째, 게이트 활성화는 softmax로 정규화해야 한다. g = softmax(s)는 log g_m = s_m − LSE(s)가 되어, 단위 게이트를 갖는 현재 블록 대비 과거 문맥의 총 어텐션 질량을 보정하고 점수의 전역 이동에 불변해진다. 시그모이드 게이트는 1로 포화되고 비정규화 로짓 주입은 0으로 붕괴해 과거 블록과 현재 블록의 구분이 흐려진다. 셋째, 순위 정보를 보존해야 한다. STE로 이진 Top-K 마스크를 쓰면 정규화 항이 선택 집합에만 걸려 탈락 블록의 가중치가 지수적으로 커지고, 초기 랜덤 선택기에서 그래디언트 노름이 소프트 게이팅보다 수십 배 커져 학습이 불안정해진다. 넷째, 학습 범위는 선택된 블록만 갱신하는 sparse scope로 충분하다. 선택되지 않은 블록은 자기 콘텐츠에 기반한 독립 그래디언트를 받지 못해 초반 수렴이 느리지만, 결국 full scope와 비슷한 최종 성능에 더 낮은 비용으로 도달한다. 구현에서는 게이트 덧셈을 FlashAttention의 타일 단위 qK^T 계산에 융합하는 Triton 커널을 직접 만들어, 전체 어텐션 행렬을 물리화하지 않고도 장문맥 학습이 가능하게 했다.

어떻게 쓰나

실험은 SeerAttention-R의 AttnGate 선택기를 그대로 쓰고, 블록 크기 64, Top-K 32, Qwen3-4B를 OpenR1-Math-220k의 93.7K 예시로 학습(학습·생성 길이 32,768, 백본 동결)해 GPQA-Diamond를 16회 평균으로 평가하는 통제 실험으로 설계됐다. 추론 과제에서 1024 토큰 예산일 때 SAS는 SeerAttention-R 대비 MATH500에서 6.0~7.7점, GPQA-Diamond에서 10.6~15.5점(Qwen3-4B/8B/14B)을 앞선다. AIME24는 예산 2048에서 +13.0(Qwen3-4B)이고, 예산 4096에서는 71.72 대 71.25로 full attention을 오히려 넘긴다. 학습 없이 쓰는 Sliding Window·StreamingLLM은 예산이 빡빡해지면 급격히 무너지고, Quest는 예산 2048에서 AIME24/25 점수가 0으로 붕괴한다. 수학 데이터로만 선택기를 학습하고 장문맥 이해(LongBench)로 바로 평가한 전이 실험에서도 SAS는 거의 모든 예산과 백본에서 SeerAttention-R을 앞서며, 입력이 길수록 격차가 벌어진다(예산 2048, Qwen3-14B의 8K+ 구간에서 53.9 대 51.5, +2.4). 예산 4096에서는 Qwen3-14B 평균 56.2 대 56.6으로 full attention을 사실상 회복한다. 에이전트 과제에서는 BFCL 멀티턴에서 Qwen3-4B 예산 2048에 +3.5, Qwen3-14B 예산 4096에서 44.00 대 44.50으로 근접하고, VitaBench에서도 예산 4096에 대부분 지표에서 앞선다.

전제와 한계

사후 학습을 넘어 continued pretraining으로도 확장한다. OLMo3-7B stage-1 체크포인트에서 블록 64, Top-K 32로, 어텐션 프로젝션에서 초기화한 전용 검색 query/key 프로젝션을 가진 경량 선택기를 만들어 백본과 함께 학습했다(OLMo3 사전학습 코퍼스, 최대 길이 8,192, 13,000 스텝, 전역 배치 512, 약 50B 토큰, AdamW 2e-4에서 2e-5로 코사인 감쇠). 일반·수학·코드 과제 평균에서 SAS는 43.28로 슬라이딩 윈도우 continued pretraining 베이스라인(43.24)과 비슷하고 HiLS-Attn-RoPE(41.68)를 앞서며 dense OLMo3-Base(43.88)에 근접한다. LongBench 평균은 30.0으로 공동 최고이며 dense 베이스 29.0, 슬라이딩 윈도우 28.0을 앞선다. 이득은 8K 초과 장입력에 집중된다.

분석 파트의 관찰이 특히 흥미롭다. SAS는 층별로 선택 블록이 덮는 어텐션 질량이 증류 방식보다 오히려 적다(가중치 qK^T 기준, 값 크기를 반영한 qK^T + log‖V‖_2 기준 모두). 대신 층 전체 선택 블록의 합집합이 full attention 오라클과 겹치는 재현율은 더 높다. 층별 국소 목표를 맞추는 증류와 달리, end-to-end 학습은 층 간 선택을 함께 최적화해 각 층의 선택이 더 상보적이 된다는 해석이다. 생성 행동에서도 Qwen3-4B 예산 4096 기준 네 개 추론 벤치마크 모두에서 SAS가 더 짧은 생성과 더 낮은 절단률을 보였고, 어려운 AIME에서 격차가 컸다. 실제 디코드 효율은 SGLang 단일 GPU, CUDA 그래프, 청크 프리필, 정상 상태 디코드(512 토큰 생성) 조건에서 측정했는데, 배치 1에서 64K 2.4배, 256K 4.6배, 512K 5.6배, 배치 8에서는 64K에서 약 13배 빨라진다. 다만 단계별 분해를 보면 어텐션 연산은 문맥 길이와 무관하게 일정한 반면, 선택기 스코어링과 Top-K 랭킹은 문맥에 비례해 커져 8K에서 21%이던 Top-K 선택이 512K에서는 90%를 차지한다.

개발자 관점에서 이 논문의 실용적 의미는 명확하다. 어텐션 커널을 게이트 어텐션 커널로 교체하고 표준 언어모델 손실만 최적화하면 되므로, 별도의 교사 모델·증류 파이프라인·층별 어텐션 통계 수집 없이 기존 파인튜닝 루프에 얹을 수 있다. 선택기 아키텍처(AttnGate)와 백본, 학습 데이터를 고정한 통제 비교에서 언어모델 손실이 층별 어텐션 증류보다 나은 지도 신호였다는 점은, 사후 희소화를 도입할 때 증류 손실을 먼저 검토할 이유가 줄어든다는 뜻이다. 다만 도입 전에 확인할 것이 있다. 첫째, 예산이 빡빡할수록 이득이 크므로 1024~2048 토큰 수준의 낮은 예산에서 이득을 측정해야 한다. 둘째, 선택기 스코어링과 Top-K 랭킹이 초장문맥에서 지배적 병목이 되므로, 512K급을 노린다면 어텐션 커널보다 선택 단계 커널을 먼저 최적화해야 한다. 셋째, SAS는 디코드 어텐션만 희소화하므로 프리필 비용은 그대로다.

한계와 전제는 논문 자체의 서술에서 확인된다. 저자들은 continued pretraining 확장을 "초기 증거(initial evidence)"로만 제시하며, 사후 학습 실험은 Qwen3 계열과 AttnGate 선택기, 수학 중심 학습 데이터(OpenR1-Math-220k)에 한정된다. sparse scope 학습은 초반 수렴이 full scope보다 느리며, 최종 성능이 비슷해진다는 것은 실험적 관찰이지 이론적 보장이 아니다. 또한 SAS는 층별 어텐션 질량 커버리지가 증류보다 낮고, 그 대신 층 간 합집합 재현율이 높다는 것이 성능 우위의 설명으로 제시되지만, 이것이 인과적으로 성능을 만든다는 직접 증명은 아니다. 논문은 별도의 한계 절을 두지 않았고, 위 내용은 본문의 분석·실험 서술에서 확인되는 전제들이다.