어텐션 스스로 계산량을 배분하는 MALA가 긴 문맥 비용을 줄인다
MassAlloc Attention: Let Attention Allocate Its Own Compute
무엇인가
긴 문맥을 다루는 풀 소프트맥스 어텐션(FullAttn)은 계산량이 시퀀스 길이의 제곱으로 늘어난다. 문제는 어텐션이 실제로 만들어내는 확률 분포가 극단적으로 비균일하다는 점이다. 논문은 FullAttn이 인과적 점수 공간의 상당 부분에 무시할 만한 정규화 질량만 배정하는데도, 밀집 커널은 각 QK 타일을 만든 뒤 소프트맥스 갱신, V 로딩, PV 누적 같은 후처리 경로 전체를 실행한다고 지적한다. 역방향에서도 확률 재구성과 dP, dS, dQ, dK, dV 계산이 모든 합법 타일에 대해 수행된다. 저자들은 어텐션이 이미 만들어낸 점수와 소프트맥스 통계가 이 결정을 내릴 정보를 갖고 있다는 점에 주목한다.
어떻게 동작하나
MALA(MassAlloc Attention)는 QK 점수 탐색과 후처리 실행을 분리한다. 모든 합법적 인과 타일에 대해 QK 점수는 그대로 계산하되, 그 이후 단계를 실행할지는 정규화 기여도로 결정한다. 기준 척도는 균등 어텐션이다. 쿼리 행 q가 볼 수 있는 인과 키의 수를 L_q라 하면 균등 어텐션은 각 키에 1/L_q를 준다. 여기에 무차원 허용치 τ를 곱해 길이 인식 임계값 τ/L_q를 만든다. τ는 계산 예산이 아니라 허용 가능한 기여도를 지정하는 값이다. 분포가 뾰족한 행은 많은 타일을 거부하고, 퍼진 행은 더 많이 유지한다. τ가 0으로 가면 건너뛰기 조건이 사라져 FullAttn으로 복원된다.
무엇과 다른가
전방 패스는 온라인 소프트맥스가 유지하는 러닝 최댓값 m_i와 이동 정규화 합 ℓ_i를 그대로 쓴다. 후보 타일의 최대 기여 비율은 exp(타일 내 최대 점수)/Z_q^(t)로 계산되고, 이것이 τ/L_q보다 작으면 타일을 건너뛴다. 로그를 취해 커널이 쓰는 형태로 만들면 rowmax(S_ij) − m_i − log ℓ_i < log(τ/L_i)이다. 아직 유지된 키가 없는 행은 로그 비율이 +∞로 정의되어 첫 합법 타일이 반드시 유지되므로 모든 전방 어텐션 행은 비어 있지 않다. 역방향은 전방이 저장한 최종 로그 정규화기 ℓ_i를 재사용해 S_ji − ℓ_i^T < log(τ/L_i)^T 조건으로 타일을 검사한다. 최종 정규화기는 온라인 중간 정규화기보다 크거나 같으므로 전방에서 건너뛴 타일은 역방향에서도 반드시 건너뛴다. 역방향은 전방이 유지한 타일 중 일부를 추가로 생략할 수 있어 유지 지원이 전방 지원에 내포되며, 전방 마스크를 저장할 필요가 없다. 다만 이것이 전방 연산자의 정확한 미분을 뜻하지는 않으며, 저자들은 그래디언트 충실도를 연산자 수준에서 직접 측정한다.
어떻게 쓰나
8K 문맥의 동일 작업량(matched-work) 연구에서는 모든 정책이 정확히 같은 총 후처리 작업량을 실행하도록 맞췄다. MALA가 τ=1에서 유도한 쿼리당 약 1,024개의 후처리 키 슬롯이 공통 기준이다. MALA의 온라인 결정은 인스턴스별 참조 질량 오라클에 근접했다. 평균 생략 질량은 0.0188% 대 0.0182%, 평균 상대 출력 L2 오차는 0.0174% 대 0.0164%였다. 정적 레이어-헤드-위치 할당과 비교하면 인스턴스별 할당이 평균 생략 질량을 9.5배, 평균 상대 출력 L2 오차를 17.2배 줄였다. 1K에서 32K까지 컨텍스트를 32배 늘려도 같은 허용치가 유지됐다. 평균 전방 작업량은 쿼리당 470에서 1,066 키 슬롯으로, 역방향은 462에서 1,058로 늘었고 8K 이후에는 둘 다 1,010~1,066 구간에 머물렀다. 평균 생략 확률 질량은 최대 0.0062%, P95는 최대 0.032%였고, 평균 상대 출력 L2 오차는 최대 0.021%(P95 0.14%), 그래디언트 오차는 dQ 0.38%, dK 0.35%, dV 0.17%(P95 각 0.90%, 0.66%, 0.23%)였다.
전제와 한계
통제된 연상 회상(associative recall) 실험에서는 고정 예산 베이스라인이 1,024 슬롯 상한을 쓰도록 맞췄다. 시퀀스 길이 8,192, d_model=512에서 MALA는 89.67%, FullAttn은 89.97%를 기록한 반면 DSA는 52.61%, MoBA는 47.25%, NSA는 22.61%에 그쳤다. 8개 H100에 텐서 병렬(TP=8)을 적용한 128K 토큰 연산자 벤치마크에서 MALA는 FullAttn 대비 학습 전방 지연을 2.2배, 역방향을 3.0배, 추론 디코딩 지연을 1.6배 줄였다. 최대 연산자 메모리는 FullAttn과 같은 수준이었고, 디코딩 지연은 DSA보다 3% 이내로 느렸다. 128개 H100에서 0.6B~14B 규모로 스케일링 법칙 학습을 돌린 결과 14B에서 퍼플렉서티 차이는 0.001 미만이면서 총 학습 FLOPs는 4K 사전학습에서 2.5%, 32K 장문맥 학습에서 23.1% 줄었다. 모델 수준 평가에서 14B MALA는 지식 72.48, 추론 64.66으로 FullAttn의 72.32, 64.46과 비슷했고, 32B에서는 75.75/76.10 대 75.62/75.67이었다. 네이티브 32K RULER는 14B에서 89.45 대 89.42, 32B에서 92.71 대 92.70이었고, YaRN으로 128K까지 외삽한 점수는 14B에서 65.75 대 65.84, 32B에서 82.56 대 82.03이었다.
실무 관점에서 MALA의 매력은 기존 어텐션 커널을 대체하는 융합 연산자라는 점이다. 선택 마스크나 유지 타일 인덱스, 라우터 상태를 물질화하지 않고 표준 어텐션 상태만 쓴다. 같은 허용치 τ가 학습 전방·역방향, 추론 프리필, 자기회귀 디코딩을 모두 관장하므로 별도 캘리브레이션이 필요 없다. 다만 QK 점수 탐색은 모든 합법 타일에 대해 그대로 수행되므로 복잡도는 여전히 제곱이고, 절감은 후처리 산술과 메모리 트래픽의 데이터 의존적 상수 배 감소다. 어텐션 분포가 퍼져 있는 워크로드에서는 유지되는 타일이 많아 이득이 줄어든다. 또한 MALA는 KV 캐시를 압축하거나 축출하지 않으므로 디코딩 메모리 관점의 이득은 연산자 작업 집합에 한정되고, 캐시 용량 문제는 별도 최적화로 남는다.
저자들이 명시한 전제와 한계는 다음과 같다. 첫째, QK 점수 탐색은 완전한 인과 문맥에 대해 유지되므로 이차 복잡도는 그대로다. 둘째, 역방향은 전방 연산자의 정확한 미분이 아니다. 역방향이 추가 그래디언트 기여를 생략할 수 있고 그 오차는 dO와 V에도 의존하기 때문에, 그래디언트 충실도는 이론이 아니라 경험적으로 평가된다. 셋째, 타일 단위 검사는 보수적이어서 한 행이라도 건너뛰기 조건을 만족하지 않으면 쿼리 블록 전체가 타일을 유지한다. 넷째, split-KV 디코딩에서는 각 분할이 자기 부분 정규화기로 검사하므로 더 보수적으로 동작할 수 있고, 공유 허용치가 분할·비분할 실행에서 동일한 지원을 유지한다는 보장은 없다. 다섯째, 실험은 특정 모델 규모(0.6B~14B, 별도 추가 학습한 32B)와 하드웨어(H100, TP=8 및 128 GPU) 구성에서 이뤄졌다.