CoWindow Attention이 KV 헤드별 분담으로 전체 인과 커버리지를 이룬다
CoWindow Attention: Full Causal Coverage Is a Collective Property
무엇인가
긴 문맥을 다루는 현재의 full attention(FullAttn)은 모든 어텐션 헤드에게 완전한 인과 프리픽스를 반복해서 노출한다. FlashAttention 계열 커널이 타일링과 online softmax로 IO 효율을 끌어올렸지만 어텐션 패턴 자체는 그대로여서, 헤드 수만큼 같은 장거리 토큰을 다시 계산하고 다시 읽는 중복이 남는다. 이 논문은 그 중복이 필수인지 묻는다. 각 헤드가 모든 과거 위치를 직접 볼 필요는 없고, 층 안의 헤드 집합이 합쳐서 모든 인과 위치를 볼 수 있으면 된다는 것이 출발점이다.
어떻게 동작하나
제안 방법인 CoWindow Attention(CoWA)은 인과 이력을 KV 헤드들에 분배한다. 각 KV 헤드에는 세 개의 윈도우가 배정된다. 모든 헤드가 공유하는 prefix-sink 윈도우(시퀀스 시작부, 폭 w_sink)와 near-diagonal 윈도우(최근 문맥, 폭 w_near), 그리고 나머지 장거리 구간을 헤드별로 하나씩 나눠 갖는 long-range 윈도우다. 거리 정의는 우하단 정렬 방식으로 δ(i_q, i_k) = i_q + (N_k − N_q) − i_k이며, 이는 학습·프리필·자기회귀 디코딩에서 동일하게 쓰이도록 만든 식이다. 장거리 구간 길이 L = max(N_k − w_sink − w_near, 0)을 H_k개로 균등 분할해 경계 b_hk = floor(h_k·L / H_k)를 두고, KV 헤드 h_k는 δ가 [w_near + b_hk, w_near + b_hk+1)에 속하는 키만 본다. GQA에서는 QO 헤드 h가 KV 헤드 floor(h·H_k/H_q)의 키·값과 윈도우 배정을 공유한다. 이 세 윈도우의 합집합이 각 쿼리에 대해 전체 인과 프리픽스와 정확히 일치하며, 저자들은 이를 collective coverage라고 부른다. 학습된 라우터나 인덱서가 필요 없고, 학습과 추론이 같은 위치 규칙을 쓰며, 전역 KV 헤드 인덱스(h_global = r·H_k_local + h_local)로 텐서 병렬 랭크 간에 상보적 윈도우가 보장된다.
무엇과 다른가
비용 분석도 명시적이다. 마지막 쿼리 위치에서 KV 헤드 전체의 가시 키 수 합은 H_k(w_sink + w_near) + L이고, dense causal은 H_k·N_k다. 즉 고정된 H_k에서 어텐션은 여전히 시퀀스 길이에 대해 이차이지만, 중복되는 쿼리-키 연결 수가 줄어드는 만큼 절약된다. H_k = 1이면 이 구조는 dense causal attention으로 퇴화한다. 실행 구조는 규칙적이다. forward/backward 커널은 QO 헤드와 블록 단위로 가시 블록 집합만 순회하고, 인과·윈도우·시퀀스 경계를 지나는 블록만 원소별 마스킹이 필요하다. 디코딩은 N_q = 1인 특수 케이스로, 키 길이가 늘 때마다 경계를 갱신해 장거리 윈도우 폭을 균등하게 유지하고 split-KV 프로그램과 online softmax 상태로 병합한다.
어떻게 쓰나
가장 직접적인 근거는 8K 윈도우 매칭 ablation이다. 헤드당 윈도우 폭을 고정한 채 서로 다른 장거리 윈도우 개수를 1·2·4·8개로 늘리면, 장거리 이력의 12.5%·25%·50%·100%를 집합적으로 커버하게 된다. 검색 정확도는 21.32% → 32.92% → 52.32% → 89.73%로 단조 증가했고, FullAttn은 89.97%, near 윈도우만 쓰는 SWA는 5.34%였다. 윈도우를 복제하는 대신 상보적으로 나누는 것 자체가 성능을 만든다는 뜻이다.
전제와 한계
통제된 associative recall 실험에서는 256개의 키-값 쌍을 학습하고 시퀀스 길이를 1,024에서 8,192까지, d_model을 64에서 512까지 늘리며 쿼리당 토큰 예산을 1,024에서 1,920으로 맞췄다. 시퀀스 8,192·d_model 512에서 CoWA는 89.73%로 FullAttn 89.97%와 사실상 동일했고, DSA 53.71%, MoBA 50.12%, NSA 25.23%, 나머지 방법들은 약 10%에 머물렀다. 모델 차원을 키우는 것만으로는 이 검색 격차가 좁혀지지 않았다.
시스템 비용은 128K 토큰, H100 8장, 텐서 병렬 TP=8 조건에서 측정됐다. CoWA는 FullAttn 대비 학습 forward 지연을 7.4배, backward를 8.6배, 추론 디코딩 지연을 3.0배 줄였다. forward·backward 피크 메모리는 FullAttn 구현과 동일했고, 디코딩 오퍼레이터 피크 할당은 8.4MiB로 FullAttn보다 7.6배 작았다. MoBA는 128K backward에서 랭크당 메모리 한도를 초과했다. MoBA의 블록 풀링·라우팅·TopK·재배열, DSA의 lightning indexer·양자화·TopK 비용과 달리 CoWA는 위치와 전역 KV 헤드 인덱스에서 가시 블록을 바로 유도한다.
0.6B에서 14B까지 128장의 H100으로 수행한 스케일링 실험에서 CoWA는 FullAttn의 perplexity를 따라가면서 총 학습 FLOPs를 줄였다. 14B 기준 4K 사전학습에서 3.1%, 32K 장문맥 학습에서 28.5% 절감이며, 장문맥 학습 후 perplexity 차이는 0.01 미만이다. 이때 QO 헤드당 토큰 예산은 4,992이고 라우터·인덱서 FLOPs는 없다. 모델 수준 평가에서도 14B 지식 72.70 대 72.32, 추론 64.87 대 64.46, 32B 지식 76.07 대 75.62, 추론 75.53 대 75.67로 비슷했고, native 32K 검색은 두 스케일 모두 0.3점 이내 차이였다. YaRN으로 128K까지 외삽한 검색은 14B에서 66.60 대 65.84, 32B에서 81.78 대 82.03이었다.
저자들이 밝힌 전제와 한계도 분명하다. 전체 인과 커버리지는 토큰 위치에 대한 직접 접근을 보장할 뿐, FullAttn과 동일한 헤드별 상호작용이나 출력을 의미하지 않는다. 각 QO 헤드는 자기 가시 키 집합 위에서 따로 softmax를 수행한다. 이 구조는 H_k > 1일 때만 성립하며, 고정된 H_k에서는 여전히 시퀀스 길이에 대해 이차 복잡도이고 절약은 중복 연결 감소에서 나온다. 학습을 위한 equal-area 할당은 이 논문의 범위 밖이고, HBM에 상주하는 KV 저장량을 줄이는 window-aware offloading·prefetching은 향후 과제로 남겨졌다. 또한 TP > H_k인 구성은 논리 KV 헤드를 복제하므로 랭크당 전역 KV 캐시의 1/TP보다 많은 양을 저장한다. 오퍼레이터 메모리 측정치는 지속적인 KV 캐시 저장량과는 별개다.
실무 관점에서 CoWA는 어텐션 커널과 텐서 병렬 배치를 크게 뜯어고치지 않고 긴 문맥 비용을 줄이려는 팀에 맞는 선택지다. 도입 전에 확인할 것은 세 가지다. 첫째, KV 헤드 수가 1보다 충분히 큰지(GQA 구성에서 KV 헤드가 1개면 이득이 없다). 둘째, 워크로드가 32K 이상 장문맥이고 TP 샤딩이 KV 헤드 단위와 정렬되는지. 셋째, 콘텐츠 기반 검색이 꼭 필요한 태스크인지다. CoWA는 위치 규칙만 쓰므로 내용에 따라 달라지는 선택이 필요한 워크로드에서는 라우터 기반 방법과 요구가 다르다.