포스트

Mixture of Sparse Attention: Expert-Choice 라우팅 기반 학습형 희소 어텐션

목차

  1. 개요
  2. 방법론
  3. 주요 결과
  4. 한계와 주의사항
  5. 결론
  6. Reference

개요

KAUST와 스탠퍼드 연구진이 발표한 Mixture of Sparse Attention(MoSA)은 Mixture of Experts(MoE)의 Expert-Choice 라우팅을 어텐션 메커니즘에 적용한 학습형 희소 어텐션 기법이다. 트랜스포머의 셀프 어텐션은 시퀀스 길이에 대해 계산량과 메모리가 모두 이차로 증가하며, 추론 시 KV 캐시의 메모리 점유가 배포의 병목으로 작용한다. State Space Model이나 선형 어텐션 같은 서브쿼드래틱 대안은 고정 크기 메모리에 과거 전체를 압축해야 하므로 실제 성능에서 완전한 셀프 어텐션에 미치지 못한다. 정적 희소 어텐션 역시 사람이 정의한 규칙 기반 패턴을 사용하기 때문에 블록 단위 손실 압축이 발생하고 세밀한 회상이 어려워진다.

MoSA는 각 어텐션 헤드를 하나의 전문가로 취급하고, 헤드가 직접 시퀀스에서 자신이 처리할 토큰 k개를 선택하도록 한다. 길이 T의 시퀀스에서 k개 토큰만 선택하면 헤드당 계산 복잡도는 O(T^2)에서 O(k^2 + T)로 줄어든다. 절약된 예산은 헤드 수를 늘리는 데 재투입되어, 동일한 FLOP 예산 안에서 훨씬 많은 수의 고도로 특화된 희소 헤드를 운용할 수 있다. 논문은 IsoFLOP 설정에서 MoSA가 최대 27% 낮은 perplexity를 달성했으며, 저자들이 분석한 희소 어텐션 기법 중 밀집(dense) 베이스라인을 능가한 유일한 방법이라고 보고한다.

방법론

토큰 선택과 라우팅

기존 MoE는 토큰이 전문가를 선택하는 방식이라 특정 전문가에 부하가 몰리는 load-balancing 문제가 발생하고, 이를 보조 손실로 완화해야 한다. Expert-Choice 라우팅은 선택 방향을 뒤집어 전문가가 자신이 처리할 입력 상위 k개를 고르며, 정의상 완벽한 부하 균형을 보장한다. MoSA는 이 구조를 어텐션 헤드에 그대로 이식한다.

각 MoSA 헤드는 표준 투영 외에 라우터 가중치 행렬 W_r을 추가로 가진다. 입력 시퀀스 X에 대해 라우터는 토큰별 선택 점수를 계산한다.

1
2
3
r = sigmoid(X * W_r)          # r ∈ R^T
r_topk, I = TopK(r, k)        # I ∈ {0, ..., T-1}^k
X_s = (X_I1, X_I2, ..., X_Ik) # X_s ∈ R^(k x h)

점수 함수로는 소프트맥스가 아닌 비경쟁적 시그모이드를 사용하며, 이는 σ-MoE의 관찰을 따른 선택이다. TopK는 상위 k개의 점수와 해당 인덱스를 반환하고, 인덱스로 원본 시퀀스에서 토큰을 gather하여 헤드 입력을 구성한다. 헤드마다 라우터가 독립적이므로 각 헤드는 서로 다른 희소 패턴을 학습하며, 블록 단위 압축 없이 개별 토큰의 정보를 그대로 보존한다.

어텐션 계산과 위치 인코딩

선택된 토큰에 대해서만 query, key, value 투영을 계산하므로 W_Q, W_K, W_V, W_O 연산 자체가 k개 토큰 분량으로 줄어든다. 인과 마스크는 선택된 부분 시퀀스가 아니라 원본 인덱스를 기준으로 구성되어야 하므로 삼각 행렬이 아니다. 구체적으로 마스크는 I_i가 I_j 이상일 때만 0이고 그 외에는 음의 무한대로 설정된다.

어텐션 출력은 라우터 점수와 곱해진 뒤 원래 위치로 되돌려진다.

1
2
3
4
A   = Attention(Q, K, V, M)
X_o = diag(r) * A * W_o        # X_o ∈ R^(k x h)
Y_j = X_o_i  (j = I_i 인 경우)
Y_j = 0      (그 외)

라우터 점수를 출력에 곱하는 연산은 토큰 기여도를 점수에 비례하게 만드는 동시에 라우터로 그래디언트가 흐르는 경로를 만든다. 이 덕분에 선택 메커니즘 전체가 언어 모델링 목적 함수로 end-to-end 학습된다. 표준 어텐션 형태를 그대로 사용하므로 Flash Attention 같은 최적화 구현과 결합할 수 있고, PyTorch에서는 einsum, scatter, gather 연산만으로 멀티헤드 버전을 구현할 수 있다.

위치 인코딩은 RoPE를 사용한다. 회전 각도가 부분 시퀀스 내 위치가 아니라 원본 시퀀스에서의 위치를 반영해야 하므로, 선택 인덱스 I를 인지하도록 RoPE를 수정했다.

하이브리드 구성도 방법론의 핵심이다. 저자들은 주 실험에서 MoSA 헤드에 밀집 헤드 4개를 항상 함께 배치했으며, 부록 실험으로 이 하이브리드 구성이 필수적임을 보였다. IsoFLOP 실험에서는 StreamingLLM의 어텐션 싱크 관찰을 반영해 모든 MoSA 헤드가 첫 번째 토큰을 항상 포함하도록 했고, 나머지 k-1개를 라우터 점수로 선택한다.

FLOP 비용 비교

h는 모델 히든 차원, h’는 헤드 히든 차원, 희소도는 γ = T / k로 정의된다. 헤드 하나당 FLOP 비용은 다음과 같다.

1
2
3
4
FLOP_dense   = 8*h*h'*T + 4*h'*T^2
FLOP_mosa    = 8*h*h'*k + 4*h'*k^2 + 2*h*T + h'*k
FLOP_fixed   = 8*h*h'*k + 4*h'*k^2
FLOP_routing = 6*h*h'*T + 4*h'*k^2*γ + 2*h*T

MoSA의 라우팅 오버헤드는 토큰 점수 계산의 2hT와 출력 스케일링의 h’k로 구성되며, 나머지 항에 비해 작다. 따라서 MoSA 헤드의 비용은 고정 희소 어텐션 헤드와 거의 같은 수준이면서 내용 기반 동적 희소성을 얻는다. Routing Transformer는 클러스터링 전에 모든 토큰의 투영을 계산해야 하므로 투영 비용이 시퀀스 전체 길이에 비례한다. FLOP 관점에서 Routing Attention 헤드 하나는 대략 γ개의 MoSA 헤드에 대응한다.

l개 레이어, 밀집 헤드 H_dense개, MoSA 헤드 H_mosa개로 구성된 모델의 순전파 총 비용은 다음과 같다.

1
2
3
l*H_dense*(8*h*h'*T + 4*h'*T^2)
  + l*H_mosa*(8*h*h'*k + 4*h'*k^2 + 2*h*T + h'*k)
  + 16*l*h^2*T

레이어 정규화, 잔차 연결, 토큰 임베딩은 밀집 모델과 MoSA 모델에 동일하게 적용되는 오버헤드이므로 FLOP 매칭 계산에서 제외했다.

비교 대상 베이스라인

Fixed Sparse Attention은 Sparse Transformer 계열의 위치 기반 정적 패턴으로, 희소도 γ에 대해 stride γ로 k = T/γ개 토큰을 선택한다. MoSA 표기법으로는 인덱스가 고정된 등간격이고 라우터 점수가 항상 1인 특수 케이스에 해당한다. 사전 선택된 토큰이 앞선 레이어에서 필요한 정보를 미리 집약하고, 이후 레이어에서 필요한 위치로 다시 전달해야 하므로 정보 라우팅 부담이 표현력을 제약한다.

Routing Transformer의 Routing Attention은 온라인 K-means로 각 헤드 안에서 토큰을 γ개 클러스터로 묶는다. MoSA와 가장 유사한 내용 기반 기법이지만, 온라인 K-means는 수렴 속도가 매우 느린 것으로 알려져 있고 키와 쿼리의 클러스터링이 언어 모델링 목적과 정렬되는지도 불분명하다. 또한 Routing Transformer는 소스와 목적지 토큰을 동일하게 맞추려면 W_Q = W_K 제약을 걸어야 하지만, MoSA는 두 투영을 분리한 채로 동일 선택을 강제할 수 있어 유연성이 크다.

실험 설정

모든 모델은 어휘 크기 8000의 SentencePiece 서브워드 토크나이저를 사용한다. 학습 데이터는 C4이며, 배치 크기 64, 시퀀스 길이 1024로 100k 배치를 학습해 약 6.5B 토큰을 소비한다. 옵티마이저는 Adam이고 학습률 0.00025, 그래디언트 클리핑 노름 0.25, 선형 워밍업 4k 스텝을 적용했다. 아키텍처는 Pre-layer normalization 트랜스포머를 기반으로 한다.

네 가지 규모의 밀집 베이스라인 하이퍼파라미터는 다음과 같다.

항목TinySmallMediumLarge
순전파 FLOPs (G)54.76219.85430.701,130.65
레이어 수691827
히든 크기5121,0241,0241,280
피드포워드 히든 크기2,0484,0964,0965,120
헤드 히든 크기64646464
헤드 수99916
파라미터 수28M113M210M516M

FLOP 매칭은 밀집 헤드를 희소 헤드로 교체하되, 희소 모델의 FLOP이 베이스라인을 넘지 않는 최대 헤드 수를 선택하는 방식으로 수행했다. 모든 희소 모델은 밀집 헤드 4개를 유지하며, 이 4개도 FLOP 계산에 포함된다.

주요 결과

IsoFLOP 언어 모델링 성능

동일 FLOP 예산에서 희소도 1을 초과하는 모든 설정 중 최고 perplexity를 정리한 결과는 다음과 같다. 괄호 안은 밀집 베이스라인 대비 상대 변화이며, perplexity는 낮을수록 좋다.

모델 규모밀집 파라미터밀집 perplexityMoSA 최고Fixed 최고Routing 최고
Tiny28M22.4616.39 (-27.0%)23.28 (+3.7%)23.33 (+3.9%)
Small113M16.0112.85 (-19.7%)16.51 (+3.1%)16.43 (+2.6%)
Medium210M13.9511.06 (-20.7%)14.35 (+2.9%)14.21 (+1.9%)
Large516M12.2010.58 (-13.3%)12.40 (+1.6%)12.24 (+0.3%)

MoSA는 네 가지 규모 전부에서 밀집 베이스라인을 유의미하게 개선했다. 반면 Fixed Sparse Attention과 Routing Attention은 모든 희소도에서 밀집 베이스라인에 도달하지 못했고, 희소도에 따른 뚜렷한 경향도 보이지 않았다.

희소도별 상세 결과

MoSA의 IsoFLOP 곡선은 희소도가 커질수록 perplexity가 단조 개선되다가 특정 지점에서 반전되는 U자 형태를 보인다. Tiny 규모에서 희소도별 하이브리드 MoSA의 perplexity와 파라미터 수, MoSA 헤드 수는 다음과 같다. 하이브리드 모델의 총 헤드 수는 표의 MoSA 헤드 수에 밀집 헤드 4개를 더한 값이다.

희소도perplexity파라미터 수MoSA 헤드 수
1 (밀집)22.4628M0
221.7634M13
420.4548M31
819.2478M69
1618.00136M142
3216.90242M276
6416.39423M505
12817.27693M848
25618.061B1277

희소도 64 부근에서 포화되고 128 이상에서 성능이 다시 나빠진다. 시퀀스 길이 1024에서 희소도 256이면 헤드당 선택 토큰이 4개에 불과해 복잡한 관계를 담기 어렵기 때문이다. 다른 규모에서는 Small이 희소도 64에서 12.85, Medium이 희소도 32에서 11.06, Large가 희소도 4에서 10.58로 최적을 기록했다. 대형 모델일수록 헤드 수 증가로 인한 메모리 요구가 커져 탐색 가능한 희소도 범위가 제한되었다.

주목할 만한 지점은 파라미터 매칭 관점에서도 이점이 나타난다는 사실이다. 희소도 8의 Medium 모델은 파라미터 442M에 perplexity 12.16으로, 파라미터 516M에 perplexity 12.20인 Large 베이스라인보다 좋다. 계산 이득을 배제하더라도 헤드의 높은 특화 자체가 성능 향상으로 이어질 수 있음을 시사한다.

자원 사용량 최적화

IsoFLOP 결과는 FLOP 매칭 상황의 품질 우위를 보여주지만, 희소 어텐션은 자원 절감 목적으로도 사용된다. 저자들은 perplexity를 밀집 베이스라인에 맞춘 뒤 벽시계 시간, GPU 메모리, KV 캐시 크기를 측정했다. 희소도는 Tiny, Small, Medium에서 32, Large에서 16으로 고정하고 MoSA 헤드 수를 늘려가며 perplexity를 맞췄다. KV 총량은 T 곱하기 밀집 헤드 수에 k 곱하기 MoSA 헤드 수를 더해 계산했다.

모델 규모구성밀집 헤드MoSA 헤드perplexity스텝당 시간 (ms)메모리 (GB)KV 총량 (K)
TinyDense9022.4613721.19.2
TinyMoSA41722.40127 (-7.3%)19.0 (-10.0%)4.5 (-51.1%)
SmallDense9016.0232632.49.2
SmallMoSA41416.01319 (-2.1%)31.4 (-3.1%)4.4 (-52.2%)
MediumDense9013.9461950.29.2
MediumMoSA41213.76592 (-4.4%)49.4 (-1.6%)4.4 (-52.2%)
LargeDense16012.20807104.116.4
LargeMoSA41612.16703 (-12.9%)94.5 (-9.2%)5.0 (-69.5%)

측정은 Tiny, Small, Medium은 A100 GPU 1장, Large는 A100 2장에서 수행했다. 전용 CUDA 커널 없이 순수 PyTorch 연산만으로 속도, 메모리, KV 캐시 세 지표를 동시에 개선했다는 점이 핵심이다. 특히 KV 캐시는 모든 규모에서 절반 이상 줄었고 Large에서는 69.5% 감소했다. 저자들은 전용 커널을 설계하면 추가적인 효율 향상이 가능할 것으로 예상한다.

긴 시퀀스로의 확장

긴 시퀀스 실험에서는 밀집 헤드 대신 로컬 어텐션을 결합했다. 긴 컨텍스트에서는 소수의 밀집 헤드도 메모리 부담이 크기 때문이며, 이는 희소 어텐션 문헌의 표준 관행이다. 모든 긴 시퀀스 모델은 6개 레이어와 히든 차원 1024를 사용한다. Routing Transformer는 로컬 헤드 4개와 Routing 헤드 4개를, Fixed와 MoSA는 희소 헤드 60개와 로컬 헤드 4개를 사용한다.

k를 64로 고정한 채 시퀀스 길이를 1024에서 8192까지 늘렸으므로 희소도는 16에서 128까지 증가한다. 헤드 수 60은 T = 1024 기준으로 대략 FLOP을 맞추기 위한 설정이며, 시퀀스가 길어질수록 MoSA와 Fixed의 비용은 Routing보다 훨씬 낮아진다. T = 8192에서 MoSA 헤드 60개의 FLOP 비용은 Routing Transformer 헤드 4개의 22.99%에 불과하다. 그럼에도 MoSA는 모든 시퀀스 길이에서 가장 낮은 perplexity를 기록했다.

다운스트림 제로샷 평가

LAMBADA, WinoGrande, BLiMP, HellaSwag, PIQA, AI2ARC 여섯 가지 벤치마크로 제로샷 성능을 측정했다. 학습 시 시퀀스 길이는 1024로 거의 일정하지만 다운스트림 입력은 훨씬 짧을 수 있다. 이를 위해 입력별로 헤드당 선택 토큰 수를 T/γ와 2 중 큰 값으로 조정했다. 각 규모와 희소 모델 유형에서는 IsoFLOP 실험 최고 perplexity를 기록한 설정을 선택했다.

모델 규모모델LAMBADAWinoGrandeBLiMPHellaSwagPIQAAI2ARC
TinyDense18.750.372.027.559.428.0
TinyMoSA25.451.964.629.159.428.6
TinyRouting14.051.366.227.857.125.9
TinyFixed17.150.672.527.758.628.1
SmallDense25.852.176.230.962.430.1
SmallMoSA30.748.562.831.860.430.2
SmallRouting19.250.770.228.057.627.3
SmallFixed24.651.675.330.163.230.2
MediumDense31.451.277.833.864.531.5
MediumMoSA27.652.275.133.965.131.6
MediumRouting10.251.565.930.357.827.8
MediumFixed29.451.477.333.064.631.5
LargeDense36.252.580.438.767.133.8
LargeMoSA32.352.877.236.665.032.2
LargeRouting27.551.176.536.264.132.5
LargeFixed32.351.779.635.966.032.2

Tiny, Small, Medium 규모에서 MoSA는 대체로 다른 모델을 앞선다. BLiMP는 예외적으로 MoSA가 일관되게 낮은 점수를 보인다. BLiMP 데이터의 상당수 예제가 10 토큰을 넘지 않을 만큼 짧기 때문이며, 희소도 64 모델이 학습 시 시퀀스의 1.56%만 선택하던 것과 달리 길이 10인 문장에서 2개 토큰은 20%에 해당해 분포 불일치가 발생한다. Large 규모에서는 perplexity 우위에도 불구하고 밀집 베이스라인이 다운스트림에서 더 좋다. 저자들은 MoE 아키텍처의 전문가 과특화 문제와, 내용 기반 희소 어텐션이 짧은 시퀀스에서 겪는 어려움을 원인으로 지목한다.

하이브리드 구성 분석

부록 실험은 밀집 헤드와의 하이브리드가 선택이 아니라 필수임을 보인다. 모든 헤드를 MoSA 헤드로 교체한 순수 MoSA 모델은 대부분의 설정에서 희소도가 높아질수록 성능이 단조 악화되었다.

희소도Tiny 하이브리드 MoSATiny 순수 MoSA
1 (밀집)22.4622.46
221.7622.96
420.4523.30
819.2424.78
1618.0029.76

Large 규모만 예외적으로 희소도 2에서 11.83으로 베이스라인 12.20보다 개선되었으나, 같은 규모 하이브리드의 10.58에는 크게 못 미친다. 순수 MoSA는 포화도 훨씬 빨라 최적 희소도가 2 또는 1에 머무른다. 학습 손실 곡선에서도 순수 MoSA 모델은 5,000에서 10,000 스텝 구간에서 학습 진전이 급격히 둔화된다. 라우터가 무작위에 가까운 초기 단계에서 어텐션 가중치가 의미 있는 패턴을 학습하지 못하고, 패턴이 없으면 라우터도 중요한 토큰을 고르지 못하는 상호 의존 문제가 원인으로 제시된다.

밀집 헤드 개수를 0에서 9까지 변화시키며 FLOP을 고정한 실험에서는 밀집 헤드 4개가 최적이었고, 이 최적값은 희소도 4와 16에서 동일하게 나타나 희소도와 무관했다. 밀집 헤드가 최소 하나는 있어야 하며, 그 이상은 수확 체감을, 4개를 넘으면 오히려 악영향을 보였다. 밀집 헤드 부재의 악영향은 희소도가 높을수록 커진다.

한계와 주의사항

토큰에 대한 top-k 선택 때문에 MoSA는 본질적으로 비자기회귀적이며, 자기회귀 시나리오에 그대로 적용하려면 별도 적응이 필요하다. 이는 MoSA만의 문제가 아니라 모든 Expert-Choice 라우팅 기법과 비자기회귀 클러스터링을 사용하는 Routing Transformer에도 공통된 제약이다. Mixture-of-Depths는 학습 후 자기회귀 분류기를 따로 학습해 특정 토큰이 선택되었을지를 예측하는 방식으로 이 문제를 우회했으며, 저자들은 이를 중요한 후속 과제로 남겼다.

perplexity 개선이 다운스트림 성능으로 항상 이어지지는 않는다. 짧은 시퀀스 과제에서 희소 어텐션이 전반적으로 약하다는 점과, MoE 계열이 언어 모델링 능력 대비 다운스트림에서 격차를 보인다는 점이 두 가지 원인이다. 실무적으로는 잘린 시퀀스로 추가 학습을 하거나 instruction tuning을 적용해 완화할 수 있다는 선행 연구가 언급된다.

실험 규모 자체도 제약이다. 가장 큰 모델이 516M 파라미터 밀집 베이스라인 기준이고, 학습 토큰은 약 6.5B에 불과하다. 긴 시퀀스 실험은 하드웨어 예산 제약으로 8192 토큰까지만 수행했으며, 저자들 스스로 이를 잠재력을 보이는 예비 분석으로 규정한다. FLOP 매칭 시 헤드 수 증가가 메모리 요구를 함께 늘리기 때문에 대형 모델에서는 탐색 가능한 희소도가 제한되었다.

Expert-Choice 라우팅은 완벽한 부하 균형을 보장하는 대신 토큰별 계산량이 달라진다. 어려운 토큰에 더 많은 연산을 배분한다는 장점이 될 수도 있으나, 일부 토큰이 과도한 연산을 받고 다른 토큰은 굶주리는 불균등 배분으로 이어질 수도 있다. 또한 현재 구현은 전용 CUDA 커널이 없는 순수 PyTorch 수준이므로 보고된 속도 이득은 최적화 여지가 남아 있는 수치다.

결론

MoSA는 어텐션 헤드를 전문가로 취급하는 Expert-Choice 라우팅을 통해 헤드별로 처리할 토큰을 학습으로 선택하고, 절약한 연산을 추가 헤드 생성에 재투입하는 구조다. 복잡도를 O(T^2)에서 O(k^2 + T)로 낮추면서, IsoFLOP 설정에서 밀집 베이스라인 대비 최대 27% perplexity 개선을 달성했다. 저자들이 비교한 희소 어텐션 기법 중 밀집 베이스라인을 넘어선 유일한 방법이며, Fixed Sparse Attention과 Routing Attention은 모든 희소도에서 베이스라인에 미치지 못했다.

perplexity 매칭 설정에서는 전용 커널 없이도 벽시계 시간, 학습 메모리, KV 캐시 크기를 동시에 줄였고 KV 캐시는 최대 69.5% 감소했다. 극단적으로 긴 시퀀스에서만 이득을 보이는 기존 희소 어텐션과 달리, 표준 길이 컨텍스트에서도 실질적인 성능 향상을 낸다는 점이 차별점이다. 다만 밀집 헤드와의 하이브리드가 학습 안정성을 위해 필수이고, 비자기회귀 선택 메커니즘과 다운스트림 전이 격차는 해결해야 할 과제로 남아 있다. 저자들은 후속 방향으로 전용 CUDA 커널 개발, MQA와 GQA 및 SwitchHead와의 결합, 다른 희소 어텐션 유형과의 조합, 비전 트랜스포머 등 타 모달리티로의 확장을 제시한다.

Reference