3차원 텐서 상태로 선형 어텐션의 기억 용량을 늘린다
Triadic Linear Attention: Three-Dimensional Recurrent States for Long-Context Sequence Modeling
무엇인가
선형 어텐션 계열 RNN은 과거 문맥을 고정 크기 상태에 압축해 상수 시간 추론을 가능하게 하지만, 그 상태 크기가 회수(recall) 능력의 상한을 정한다. 키와 값을 외적(outer product)해 d×d 행렬 상태를 만드는 기존 선형 어텐션은 서로 직교하는 키를 최대 d개까지만 구분할 수 있고, 상태를 키우려면 헤드를 늘리거나 값 차원을 키우는 방식이어서 파라미터가 함께 커진다. 이 논문은 상태를 키우는 값싼 축을 하나 더 찾는 문제를 다룬다.
어떻게 동작하나
제안하는 triadic linear attention은 키 k, 두 번째 키 k′, 값 v 세 벡터의 3중 외적을 d×E×d 크기의 3차 텐서 상태 S에 더해 쓰고, 두 개의 질의 q, q′로 두 키 축을 각각 수축(contraction)해 읽는다. 출력은 Σ_{s≤t} (qᵀk_s)(q′ᵀk′_s) v_s 형태다. E차원 두 번째 키·질의를 만드는 두 개의 투영만 추가하면 상태가 E배로 커지므로 상태 크기 대비 파라미터 비율이 점근적으로 좋아진다. E=1이고 k′=q′=1이면 기존 선형 어텐션으로 정확히 환원된다.
무엇과 다른가
데이터 의존적 망각은 두 번째 키 축을 따라 각 슬라이스 S[:,e,:]에 별도 감쇠 게이트 α_{t,e}를 주는 방식으로 확장된다. 델타 규칙은 두 키로 읽어낸 현재 저장값 S_{t-1}×₁k_t×₂k′_t를 먼저 지우고 새 값 쪽으로 보간하는 형태가 된다. 두 키의 외적을 d·E 차원 결합 키 κ=k⊗k′로 펼치면 이 갱신은 키 차원이 d·E인 Gated DeltaNet과 동일한 형태가 된다. 청크 병렬 학습에서는 결합 키의 크로네커 구조 덕분에 (q⊗q′)ᵀ(k⊗k′)=(qᵀk)(q′ᵀk′)로 분해되어 마스크 어텐션이 (QKᵀ)⊙tril(Q′K′ᵀ)로 계산되고, 비용이 C²(d·E)가 아니라 C²(d+E)로 줄어든다. 상태는 값 축을 따라 32열 블록으로 쪼개 스레드 블록마다 배치해, E=8일 때 한 헤드 상태 512KiB(FP32)를 한 스트리밍 멀티프로세서가 통째로 들고 있지 않게 만든다.
어떻게 쓰나
용량 검증은 MQAR(multi-query associative recall)로 한다. 2층 4헤드, d=16의 최소 구성에서 E를 2·4·8·16으로 키우면 정확도 곡선이 N 기준으로 대략 두 배씩 오른쪽으로 이동했고, E=16은 기존 선형 어텐션보다 약 16배 많은 키-값 연관을 저장했다. 이때 비임베딩 파라미터 증가는 1.08배에 그쳤다.
전제와 한계
본 실험은 400M(24층, d_model 1024, 8헤드)과 1.3B(24층, d_model 2048, 16헤드), 헤드 차원 d=128로 진행했고, 베이스 믹서는 Gated DeltaNet(GDN)과 sGLA(GDN에서 델타 규칙의 erase 항을 뺀 것)다. E=8은 파라미터를 1.2%만 늘린다. 50토큰/파라미터(Chinchilla 최적의 2.5배, 400M은 20B 토큰, 1.3B은 65B 토큰)로 Fineweb-Edu에서 4k 문맥으로 사전학습한 뒤 64k로 문맥 확장했다. GDN은 약 10k 문맥부터 Transformer에 뒤지지만 E=2는 그 교차점을 크게 밀어내고, E=8은 64k에서도 Transformer보다 다음 토큰을 잘 예측한다. 회수 집약 과제에서는 두 스케일 모두 성능이 크게 올랐고 특히 긴 문서에서 정보를 복사해야 하는 FDA와 SWDE에서 이득이 컸다. RULER의 needle-in-a-haystack에서도 상태가 클수록 더 긴 문맥까지 검색이 정확했다.
동일 상태 크기 비교에서 triadic GDN은 2배·4배 조건 모두 WikiText와 모든 PG19 구간에서 가장 낮은 퍼플렉시티를 냈고 회수 벤치마크 평균도 가장 높았다. 대안들(큰 헤드, 넓은 값, 그룹 값, 헤드 추가)은 2배에서는 GDN 대비 거의 개선이 없었고 4배에서는 PG19 모든 구간에서 오히려 나빠졌으며, 그중 둘은 파라미터를 맞추려고 MLP 폭을 2816에서 640으로 줄여야 했다. 사전학습된 E=1 모델을 E=8로 확장(망각 게이트 복사, 두 번째 키·질의 투영은 새로 초기화)한 뒤 문맥 확장만 해도, 처음부터 triadic으로 학습한 모델이 얻는 이득의 절반에서 3/4가량을 회복했다. GDN/GQA 하이브리드에서는 3:1 Triadic GDN/GQA-8(E=4)이 모든 문맥 구간에서 가장 낮은 퍼플렉시티와 최고의 제로샷·NIAH 평균을 냈고, KV 캐시를 2배로 키운 3:1 GDN/GQA-4는 약 4.6k 토큰부터 더 많은 메모리를, 64k에서는 거의 두 배를 요구했다.
1.3B 모델, 배치 4, 단일 H100에서 측정한 블록당 순전·역전 시간에서 Triadic GDN은 E=2·E=4는 4k, E=8은 8k부터 Transformer를 앞선다. E=8은 4k에서 3% 느리지만 그보다 긴 모든 문맥에서 빠르고 64k에서는 5.1배 빠르다. 기존 GDN 대비 오버헤드는 E=8이 28~30%, E=4가 14~15%, E=2가 9~11%다. 절제 실험에서 결합 키 차원 d_k·E=1024를 고정하고 d_k와 E를 비슷하게 맞출수록 퍼플렉시티가 조금 나빠졌고, 두 번째 키·질의 활성화 함수는 softplus·sigmoid 같은 비음수 함수가 SiLU나 무활성화보다 일관되게 좋았다. 음수 값이 허용되면 두 번째 질의가 슬라이스를 합칠 때 기여가 상쇄될 수 있다는 설명이다.
실무적으로는 긴 문맥·회수 집약 작업에서 선형 어텐션 블록의 상태를 키우고 싶을 때, 헤드나 값 차원을 늘리는 대신 두 번째 키 축을 추가하는 선택지가 생긴다는 뜻이다. 다만 상태가 커진 만큼 학습이 느려진다. 저자들은 최적화된 커널을 써도 E=4에서 약 15%, E=8에서 약 30% 학습이 느려진다고 밝혔다. 일부 회수 집약 과제는 여전히 Transformer에 뒤지며, 그마저도 긴 문맥에서 훨씬 큰 상태를 쓸 때의 비교다. 또한 triadic 구성을 GDN과 sGLA 두 가지에만 적용했으므로 행렬 상태를 쓰는 다른 선형 어텐션 변형으로의 확장은 아직 검증되지 않았다.