지난 글에서 인코더-디코더를 노이즈 복원으로 사전학습해 번역·요약을 한 틀에 담는 방법을 봤다. 이번에는 구조가 아니라 비용을 본다. Self-Attention은 시퀀스 길이 에 대해 시간과 메모리를 모두 으로 쓴다. 문장 몇 개를 다루던 시절에는 보이지 않던 항이지만, 긴 문서와 코드베이스를 통째로 넣는 지금은 이 항 하나가 모델이 다룰 수 있는 길이를 정한다.
이 문제를 푸는 방법은 여럿이지만, 세는 방식을 바꾸면 크게 둘이다. 하나는 참조 관계를 잘라 내거나 어텐션 자체를 다른 연산으로 갈아 끼워 계산량을 줄이는 길이고, 다른 하나는 계산은 한 톨도 빼지 않은 채 GPU 안에서 데이터가 오가는 거리만 줄이는 길이다. 앞쪽은 모델이 내놓는 값이 달라지고 뒤쪽은 달라지지 않는다. 이 차이가 어느 쪽을 언제 쓸지를 거의 다 정한다.
이차 복잡도
어텐션 행렬의 크기
병목의 출처는 한 줄이다.
import torch
def naive_attention(Q, K, V): # Q, K, V: (batch, heads, N, d)
scores = Q @ K.transpose(-2, -1) # (batch, heads, N, N) ← N² 원소가 여기서 생긴다
weights = torch.softmax(scores / Q.size(-1) ** 0.5, dim=-1)
return weights @ V
의 결과는 행렬이다. 모든 쿼리가 모든 키와 내적을 하므로 셀이 개 생기고, 이 행렬은 softmax를 지나 값 행렬과 곱해질 때까지 어딘가에 놓여 있어야 한다. 이면 262,144개, 이면 16,777,216개다. 길이를 여덟 배로 늘렸는데 셀은 예순네 배가 됐다.
숫자에 단위를 붙이면 체감이 달라진다. 반정밀도(fp16)로 한 셀이 2바이트이므로 짜리 어텐션 행렬 하나가 33.6MB다. 헤드가 32개면 한 층에서만 1.07GB이고, 배치를 8로 잡으면 8.6GB다. 그리고 softmax를 통과한 확률 행렬이 같은 크기로 하나 더 필요하다. 층 하나가 아니라 어텐션 하나의 이야기다.
여기서 중요한 것은 증가의 모양이다. 컨텍스트를 4k에서 8k로 늘리는 결정은 메모리를 두 배 쓰겠다는 결정이 아니라 네 배 쓰겠다는 결정이다. 파라미터 수나 층 수를 늘릴 때의 선형 증가에 익숙한 감각으로 접근하면 매번 예산을 두 배로 틀린다.
연산량과 메모리의 분리
같은 이라도 둘은 다른 자원이다. 연산량은 이고 — 셀마다 차원 내적을 하므로 헤드 차원이 곱해진다 — 메모리는 이다. 이 둘이 갈린다는 점을 잡아 두면 뒤에 나올 방법들이 서로 무엇을 줄이는지가 한눈에 정리된다.
희소 어텐션과 슬라이딩 윈도는 셀 자체를 지운다. 계산하지 않는 셀은 저장할 필요도 없으니 둘 다 줄어든다. 상태 공간 모델은 어텐션 행렬을 아예 만들지 않으므로 둘 다 사라진다. 반면 FlashAttention은 연산량을 한 톨도 줄이지 않는다. 모든 셀을 다 계산한다. 줄이는 것은 그 셀들이 GPU 메모리를 오가는 횟수뿐이다.
그래서 복잡도 표기만 보면 FlashAttention은 아무것도 안 한 것처럼 보인다. 그대로다. 실측 속도가 몇 배씩 뛰는 이유는 복잡도 표기가 세지 않는 상수 — 메모리 이동 — 가 실제 실행 시간의 대부분을 차지하고 있었기 때문이다.
근사와 정확 계산
두 갈래를 가르는 질문은 하나다. 같은 입력을 넣었을 때 나오는 숫자가 달라지는가.
희소 어텐션·슬라이딩 윈도·SSM은 달라진다. 참조하지 않기로 한 자리는 어떤 값이었어도 결과에 반영되지 않으므로, 이들은 모델 구조를 바꾸는 결정이다. 학습을 다시 하거나 최소한 그 구조로 학습된 체크포인트가 있어야 하고, 잘라 낸 자리가 그 태스크에서 실제로 중요했는지는 돌려 보기 전에는 모른다.
FlashAttention은 달라지지 않는다. 부동소수점 누적 순서가 바뀌어 마지막 자릿수에서 미세한 차이가 생기는 정도이고, 수학적으로는 같은 값을 계산한다. 그래서 이미 학습된 모델에 그대로 얹을 수 있고 품질을 재확인할 필요도 없다. 드롭인이라는 성질이 이 방법을 표준으로 만들었다.
희소 어텐션
지역 윈도와 글로벌 토큰
희소 어텐션(Sparse Attention)은 「모든 토큰이 모든 토큰을 볼 필요는 없다」는 가정에서 출발한다. 문장 안의 대명사는 대개 근처의 명사를 가리키고, 코드의 변수는 대개 같은 함수 안에서 쓰인다. 전체 참조 중 실제로 쓰이는 몫이 얼마 안 된다면, 미리 정한 패턴만 남기고 나머지 셀은 계산하지 않는다.
2020년 Allen AI가 발표한 Longformer는 이 패턴을 둘로 나눠 세웠다. 하나는 지역 윈도로, 각 토큰이 자기 앞뒤 개만 참조한다. 셀이 토큰마다 개뿐이므로 다. 다른 하나는 글로벌 토큰으로, [CLS]나 [SEP] 같은 소수의 자리만 전체를 보고 전체에게도 보인다. 그런 자리가 개면 이고, 와 는 길이와 무관한 상수이므로 합쳐도 이다.
숫자를 넣어 보면 절약이 분명하다. , 윈도 크기 512라면 — 구현에서 이 값은 양쪽을 합친 폭이라 실제로는 좌우 256개씩이다 — 셀은 개로 전체 참조의 8분의 1이다. 길이를 8192로 늘리면 전체 참조는 넷을 곱한 값이 되지만 지역 윈도는 둘을 곱한 값이 된다. 늘릴수록 격차가 벌어진다.
글로벌 토큰을 따로 둔 이유는 분류 헤드에 있다. 문서 분류는 [CLS] 자리의 벡터 하나를 읽어 답을 내므로, 그 자리만은 문서 전체를 봐야 한다. 지역 윈도만 남기면 [CLS]는 자기 옆 256개 토큰만 요약한 벡터가 되고, 뒤쪽 3,800개 토큰은 답에 닿지 못한다. 패턴을 자를 때 무엇이 잘리는지는 태스크가 정한다.
BigBird
같은 해 Google이 낸 BigBird는 Longformer의 지역·글로벌에 랜덤 어텐션을 하나 더 얹었다. 토큰마다 무작위로 고른 몇 개를 추가로 참조하게 하는 방식이다.
무작위가 도움이 되는 이유는 참조 관계를 그래프로 보면 나온다. 지역 윈도만 있는 그래프에서 양 끝의 두 토큰이 이어지려면 중간의 모든 윈도를 하나씩 건너야 하고, 그 홉 수가 층 수보다 많으면 정보가 아예 닿지 못한다. 여기에 임의 간선 몇 개를 뿌리면 그래프의 지름이 급격히 줄어든다 — 아는 사람 몇 다리만 건너면 누구에게든 닿는다는 그 성질이다. 글로벌 토큰이 지정된 허브를 만드는 방식이라면 랜덤 어텐션은 지름길을 흩뿌리는 방식이다.
BigBird 논문은 이 세 패턴을 합친 모델이 튜링 완전(Turing complete)함을 증명해, 전체 어텐션이 표현할 수 있는 것을 이 구조도 표현할 수 있다고 보였다. 다만 증명이 말해 주는 것은 표현력의 상한이다. 같은 데이터와 같은 학습 예산에서 같은 성능이 난다는 뜻은 아니고, 실제로 무엇이 필요한지는 태스크마다 확인해야 한다.
슬라이딩 윈도와 수용 영역
슬라이딩 윈도는 희소 패턴 중 가장 단순한 것만 남긴 형태다. 글로벌도 랜덤도 없이, 각 토큰이 최근 개만 참조한다. Mistral 7B가 모든 층을 이 방식으로 세웠고, Gemma 2·3은 지역 층 사이에 전체를 보는 층을 끼워 넣는 쪽을 골랐다 — Gemma 3은 윈도 1,024짜리 지역 층 다섯에 전역 층 하나를 되풀이한다.
한 층만 보면 밖은 못 본다는 뜻이지만, 층을 쌓으면 사정이 달라진다. 두 번째 층의 한 토큰은 첫 번째 층 출력 개를 보고, 그 각각은 이미 자기 앞 개를 요약한 값이다. 층을 올라갈수록 수용 영역(receptive field) — 한 위치의 출력에 영향을 줄 수 있는 입력 범위 — 이 씩 넓어져, 층이면 이론상 까지 닿는다. 인 32층 모델이면 13만 토큰이다. 합성곱 신경망에서 작은 커널을 쌓아 넓은 영역을 보는 것과 같은 원리다.
물론 「닿는다」와 「제대로 전달된다」는 다르다. 먼 정보는 여러 층의 요약을 거치며 다른 정보와 섞이므로, 직접 참조하는 것만큼 선명하지 않다. 정확한 값 하나를 멀리서 그대로 끌어와야 하는 과제에서는 이 방식이 약하다.
실전에서 이 구조를 고르는 이유는 정확도보다 추론 메모리 쪽이다. 자기회귀 생성에서는 이미 만든 토큰의 키·값을 KV 캐시에 쌓아 두는데, 전체 어텐션이면 이 캐시가 생성 길이에 비례해 계속 자란다. 슬라이딩 윈도는 밖의 항목을 다시 볼 일이 없으므로 캐시를 로 고정할 수 있다. 컨텍스트가 아무리 길어져도 추론 한 건이 쓰는 메모리가 그대로다. 서버가 동시에 물 수 있는 요청 수가 여기서 정해진다.
상태 공간 모델
선형 재귀
지금까지가 어텐션의 셀을 골라내는 방법이었다면, 상태 공간 모델(State Space Model, SSM)은 어텐션을 쓰지 않는다. 대신 고정 크기 상태 하나를 들고 한 스텝씩 갱신한다.
순환 신경망과 모양이 같다. 결정적으로 다른 점은 갱신식 안에 비선형 함수가 없다는 것이다. 선형 재귀는 결합법칙이 성립하므로 병렬 스캔으로 접을 수 있고, 그래서 학습할 때 시간 축을 순서대로 밟지 않아도 된다. RNN이 병렬화를 못 해 트랜스포머에 밀렸던 자리를, 비선형을 포기하는 대가로 되찾은 셈이다.
비용 구조가 어텐션과 근본적으로 다르다. 시퀀스를 한 번 훑는 데 이고, 추론 중에 들고 있어야 하는 것은 상태 하나뿐이라 이다. KV 캐시라는 개념 자체가 없다 — 지나간 토큰을 저장하지 않고 상태에 접어 넣기 때문이다.
선택적 SSM
, , 가 학습된 상수이면 문제가 하나 생긴다. 무엇을 오래 기억하고 무엇을 흘려보낼지가 입력과 무관하게 정해진다는 것이다. 어텐션은 지금 쿼리에 맞는 자리를 골라 볼 수 있는데, 상수 파라미터 SSM에는 그 선택성이 없다. 같은 감쇠율로 모든 입력을 똑같이 접는다.
2023년 발표된 Mamba는 이 지점을 고쳤다. 선택적 SSM(selective SSM)은 , 와 스텝 크기를 현재 입력 에서 계산해 매 스텝 다르게 만든다. 지금 들어온 토큰이 중요하면 상태에 크게 반영하고, 아니면 거의 흘려보낸다. 게이트가 달린 순환 셀이 하던 일과 문제의식이 같고, 그것을 선형 재귀 틀 안에서 구현한 것이다.
대가가 없지는 않다. 파라미터가 시간에 따라 변하면 재귀 전체를 하나의 합성곱으로 미리 접어 두는 기존 최적화가 통하지 않는다. Mamba는 그래서 GPU 메모리 계층을 직접 겨냥한 병렬 스캔 커널을 따로 짰다 — 뒤에 볼 FlashAttention과 정확히 같은 발상이 여기서도 쓰였다.
기억 압축의 대가
| 특성 | Self-Attention | SSM (Mamba) |
|---|---|---|
| 시퀀스 길이에 대한 복잡도 | ||
| 병렬 학습 | 그대로 가능 | 병렬 스캔 커널 필요 |
| 추론 중 상태 크기 | (KV 캐시) | |
| 지나간 토큰 되짚기 | 원본 그대로 | 상태에 섞여 있음 |
표의 마지막 줄이 이 구조의 성질을 요약한다. 어텐션이 의 KV 캐시를 지고 다니는 것은 낭비가 아니라 과거를 원본 그대로 들고 있다는 뜻이다. 그래서 3천 토큰 앞에 나온 계좌번호를 지금 정확히 복사해 올 수 있다. 고정 크기 상태로 접어 넣은 쪽은 그 값을 다른 정보와 섞어 저장했으므로 같은 일을 하기 어렵다. 긴 문맥을 요약하거나 흐름을 잇는 일에는 강하고, 정확한 검색과 복사에는 약하다.
그래서 실무에서 시도되는 방향은 양자택일이 아니라 혼합이다. 대부분의 층은 SSM으로 두어 길이에 대한 비용을 선형으로 낮추고, 몇 개 층만 어텐션으로 남겨 정확한 되짚기를 담당시킨다. 어느 쪽이 얼마나 필요한지는 다루려는 과제가 정한다.
FlashAttention의 IO 최적화
메모리 이동 병목
지금까지는 계산을 줄이는 이야기였다. FlashAttention은 반대쪽에서 온다. 2022년 Tri Dao 등이 발표한 이 방법은 어텐션의 알고리즘을 바꾸지 않고, 오직 데이터가 GPU 안에서 오가는 방식만 바꿔 어텐션 연산을 7.6배까지 빠르게 만들었다.
출발점은 GPU 메모리가 한 덩어리가 아니라는 사실이다. HBM(High Bandwidth Memory)은 GPU 보드에 붙은 메인 메모리로, A100 기준 40~80GB에 대역폭이 초당 2TB쯤이다. SRAM은 연산 유닛 바로 옆에 있는 온칩 캐시로, A100은 SM 108개가 각각 192KB씩 갖고 있어 다 합쳐도 20MB 남짓이지만 대역폭이 초당 19TB에 이른다. 열 배 빠른 대신 수천 배 작다.
표준 어텐션이 이 계층을 어떻게 쓰는지 따라가 보면 문제가 보인다.
| 단계 | HBM 동작 | 크기 |
|---|---|---|
| Q, K, V 읽기 | read | |
| 쓰기 | write | |
| S 읽어 softmax | read | |
| 쓰기 | write | |
| P, V 읽어 | read |
짜리 왕복이 네 번이다. 앞에서 센 값을 그대로 쓰면 , 헤드 32개, 배치 8일 때 행렬 한 벌이 8.6GB이므로 네 번이면 34GB가 오간다. 대역폭 2TB/s로 나누면 17ms다. 그동안 연산 유닛은 데이터를 기다리며 놀고 있다. 어텐션이 이론 FLOPS 대비 한참 느렸던 이유가 여기 있었다 — 계산이 무거운 것이 아니라 계산할 것을 가져오는 데 시간을 다 썼다.
타일링
해법의 첫 번째 축은 타일링(tiling)이다. Q, K, V를 SRAM에 올라갈 크기의 블록으로 자르고, 블록 한 쌍에 대한 어텐션을 SRAM 안에서 끝까지 계산한다. 점수 행렬을 만들고, softmax를 씌우고, 값 행렬과 곱해 출력에 누적하는 데까지가 칩 위에서 일어난다. HBM으로 내려가는 것은 최종 출력 한 번뿐이다.
블록 크기는 SRAM 용량이 정한다. Q 블록 하나와 K·V 블록 하나가 중간 계산 공간까지 포함해 동시에 올라가야 하므로, 헤드 차원이 클수록 블록 행 수는 작아진다. 커널이 GPU 세대마다 다시 튜닝되는 이유가 이것이다 — SRAM 크기가 바뀌면 최적 블록 모양이 바뀐다.
행렬이 한 번도 통째로 존재하지 않으므로 메모리는 이다. 블록 쌍을 빠짐없이 도는 것은 그대로다.
Online Softmax의 누적
타일링에 걸림돌이 하나 있다. softmax는 행 전체를 봐야 계산되는 연산이다. 수치 안정을 위해 각 행에서 최댓값 을 빼고 지수를 취하는데, 블록을 하나씩 처리하는 중에는 그 행의 최댓값이 얼마일지 아직 모른다.
Milakov와 Gimelshein이 2018년에 제안한 online softmax가 이 자리를 푼다. 지금까지 본 최댓값 과 분모 누적값 을 들고 다니다가, 새 블록에서 더 큰 값이 나오면 기존 누적을 새 기준으로 다시 눈금을 맞춘다. 이전 값들이 로 쌓여 있었으니 전체에 를 곱하면 가 된다. 곱셈 한 번으로 과거를 소급 보정하는 것이다.
숫자로 한 번 따라가면 간단하다. 첫 블록의 최댓값이 3이고 분모 누적이 10이라 하자. 다음 블록에서 5가 나오면 기준을 5로 올리고, 기존 10에 를 곱해 1.35로 줄인 뒤 새 블록의 기여를 더한다. 출력 누적 도 같은 계수로 함께 보정한다.
# Online Softmax 누적 원리 (의사코드)
m_i, l_i, O_i = -inf, 0, 0 # 최댓값 · 분모 · 출력 누적
for block_j in range(num_blocks):
S_ij = Q_i @ K_j.T / sqrt(d) # 타일 어텐션 점수
m_new = max(m_i, S_ij.max()) # 기준 갱신
rescale = exp(m_i - m_new) # 과거를 새 기준으로
l_new = rescale * l_i + exp(S_ij - m_new).sum()
O_i = (rescale * l_i * O_i + exp(S_ij - m_new) @ V_j) / l_new
m_i, l_i = m_new, l_new
여기서 꼭 짚어야 할 것은 이 결과가 근사가 아니라는 점이다. 모든 블록을 돈 뒤의 과 은 행 전체를 한 번에 본 값과 정확히 같고, 따라서 출력도 같다. 부동소수점 덧셈 순서가 달라 마지막 자릿수가 흔들릴 수 있는 정도이며, 이는 배치 크기를 바꿨을 때 생기는 차이와 같은 종류다.
HBM 접근량은 표준 구현의 에서 으로 줄어든다. 은 SRAM 크기다. 헤드 차원 가 64~128이고 이 수십만 바이트인 실제 설정에서는 이 보다 한참 작으므로, 그 비율만큼 왕복이 사라진다. 논문이 A100에서 잰 값으로는 어텐션 연산만 놓고 7.6배, 학습 전체로는 GPT-2에서 3배, Long Range Arena에서 2.4배였다.
역전파에서의 재계산
학습에는 문제가 하나 더 있다. 역전파에서 어텐션 가중치 가 필요한데, 순전파에서 그것을 저장하지 않았다면 다시 만들어야 한다. 표준 구현은 그래서 짜리 를 통째로 보관한다.
FlashAttention은 대신 블록마다 소프트맥스 통계값만 남긴다. 최댓값 과 분모 이면 되고, 이 둘은 행마다 스칼라 하나씩이라 이다. 역전파에서 필요한 블록의 점수 행렬을 다시 계산한 뒤 저장해 둔 통계값을 적용하면 원래의 블록이 그대로 복원된다.
연산은 늘어난다. 순전파의 점수 계산을 한 번 더 하는 셈이다. 그럼에도 전체가 빨라지는 이유는 앞 절에서 본 그대로다 — 이 워크로드에서는 IO가 연산보다 훨씬 비싸다. 다시 계산하는 비용보다 HBM 왕복을 없앤 이득이 크다. 활성값을 버렸다가 다시 만드는 그래디언트 체크포인팅과 발상이 같고, 다만 여기서는 재계산 결과가 HBM에 내려가지 않고 SRAM 안에서 소비된다는 점이 다르다.
실질적인 효과는 학습 설정에서 나타난다. 저장할 것이 에서 으로 줄었으므로, 같은 GPU에서 배치를 키우거나 시퀀스를 길게 잡을 수 있다. 속도 이득보다 이쪽이 더 큰 변화인 경우도 많다.
커널 적용
PyTorch 내장 커널
PyTorch 2.0부터 torch.nn.functional.scaled_dot_product_attention이 이 커널을 품고 있다. 별도 설치 없이, 조건이 맞으면 자동으로 선택된다.
import torch.nn.functional as F
# 조건이 맞으면 FlashAttention 커널이 자동으로 선택된다
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
# q, k, v: (batch, heads, N, head_dim)
내부에는 백엔드가 여럿 있다. FlashAttention 커널, 메모리 효율 커널, 그리고 어디서나 도는 기본 구현이다. 조건이 맞지 않으면 아래쪽으로 조용히 내려간다. 어떤 백엔드를 쓸지 강제하는 컨텍스트 매니저도 있다. 지금 문서가 권하는 것은 torch.nn.attention.sdpa_kernel이고 예전에는 torch.backends.cuda 아래에 있었으므로, 쓰는 버전의 문서를 확인하는 편이 안전하다. torch.compile로 모델을 감싸면 이 선택이 다른 최적화와 함께 이뤄진다.
torch.nn.MultiheadAttention처럼 상위 모듈을 쓰고 있다면 내부에서 같은 함수를 부르므로 코드를 바꿀 것이 없다. 직접 짠 어텐션에서 softmax(QK^T/√d) @ V를 손으로 적어 두었다면 그 세 줄을 위 한 줄로 바꾸는 것이 이 글에서 가장 값싼 개선이다.
폴백 조건
문제는 폴백이 조용하다는 것이다. 오류도 경고도 없이 느린 경로로 내려가므로, 커널이 붙었다고 믿은 채 몇 배 느린 학습을 돌리는 일이 생긴다.
가장 흔한 원인은 헤드 차원이다. CUDA 커널은 특정 헤드 차원에 맞춰 손으로 짜여 있어서, 2의 거듭제곱에서 벗어난 값은 GPU 세대에 따라 아예 커널이 없을 수 있다. PyTorch 2.0 시절에는 헤드 차원이 80이나 96인 모델이 A100이 아닌 GPU에서 이 이유로 폴백했다. 모델 설계에서 은닉 크기를 헤드 수로 나눈 값이 어중간하게 떨어지면 여기에 걸린다. 하드웨어도 조건이다 — 이 커널은 CUDA GPU를 겨냥해 작성되었으므로 CPU나 비CUDA 가속기에서는 붙지 않는다.
마스크도 자주 걸리는 자리다. 인과 마스크를 큰 float 텐서로 만들어 attn_mask로 넘기는 것보다 is_causal=True로 알려 주는 편이 낫다. 앞쪽은 임의의 마스크로 취급되어 일반 경로를 타기 쉽지만, 뒤쪽은 커널이 블록 단위로 「이 블록은 통째로 마스킹된다」를 판단해 계산 자체를 건너뛸 수 있다. 정보를 텐서로 넘기느냐 플래그로 넘기느냐의 차이가 실행 경로를 가른다.
확인하는 방법은 결국 재는 것이다. 같은 입력으로 시간을 재 봤을 때 길이를 두 배로 늘렸는데 시간이 네 배가 된다면 폴백을 의심할 자리다.
v1에서 v3까지
| 버전 | 주요 개선 | 속도 향상 |
|---|---|---|
| v1 (2022) | IO-Aware 타일링 원형 | 표준 대비 어텐션 연산 7.6배 |
| v2 (2023) | 병렬화 개선, GQA 지원 | v1 대비 약 2배 |
| v3 (2024) | H100의 TMA와 FP8 지원 | v2 대비 1.5~2배 |
v2는 알고리즘이 아니라 일감을 나누는 방식을 고쳤다. 쿼리 블록 쪽으로도 병렬화를 열어 GPU의 워프 활용률을 높인 것이 핵심이다. v3는 H100의 Tensor Memory Accelerator(TMA)와 비동기 파이프라인을 써서, 데이터를 옮기는 동안 연산을 멈추지 않도록 겹쳤다.
버전이 GPU 세대에 묶여 있다는 점을 눈여겨볼 만하다. 커널이 특정 하드웨어의 SRAM 크기와 전송 유닛을 직접 겨냥해 쓰였기 때문이다. 알고리즘이 하드웨어와 함께 진화한다는 뜻이고, 새 GPU를 도입할 때 커널 버전도 함께 확인해야 한다는 실무적 함의가 따라온다.
다중 GPU 확장
FlashAttention이 풀지 못하는 것도 분명하다. 메모리를 으로 줄였을 뿐 0으로 만든 것은 아니므로, 시퀀스가 충분히 길어지면 한 장의 GPU에 Q·K·V와 출력을 올리는 것 자체가 불가능해진다.
이 벽은 커널이 아니라 분산으로 넘는다. 시퀀스 병렬(sequence parallelism)은 시퀀스를 여러 GPU에 쪼개 나눠 갖는 방식이고, Ring Attention은 각 GPU가 자기 몫의 Q를 들고 K·V 블록을 고리 모양으로 주고받으며 어텐션을 완성한다. Ulysses 계열은 헤드 축과 시퀀스 축 사이에서 데이터를 재배치해 통신량을 줄인다. 어느 쪽이든 각 GPU 안에서 도는 것은 여전히 FlashAttention 커널이다. 타일링이 칩 안에서 하던 일을 노드 사이에서 한 번 더 하는 셈이라, 둘은 경쟁 관계가 아니라 층이 다른 같은 아이디어다.
선택 기준
방법별 비교
| 방법 | 연산 복잡도 | 메모리 | 결과가 달라지는가 | 채택 모델 |
|---|---|---|---|---|
| 전체 어텐션 | 기준 | BERT, GPT | ||
| 희소 어텐션 | 달라짐 | Longformer, BigBird | ||
| 슬라이딩 윈도 | 달라짐 | Mistral, Gemma | ||
| SSM | 완전히 다른 구조 | Mamba | ||
| FlashAttention | 같음 | 사실상 전부 |
표의 넷째 열이 결정의 대부분을 정한다. 이미 학습된 모델의 학습·추론을 빠르게 하고 싶다면 답은 FlashAttention 하나다. 다른 선택지가 없다시피 하고, 그래서 오늘날 대부분의 프레임워크가 기본으로 켜 둔다. 반대로 처음부터 아주 긴 컨텍스트를 겨냥해 모델을 설계하는 자리라면 위쪽 셋이 후보에 오른다. 잘라 낼 참조가 태스크에서 실제로 덜 중요한지가 판단 기준이고, 이것은 돌려 보고 정하는 문제다.
복잡도 표기만 놓고 고르면 틀리기 쉽다는 점도 다시 짚어 둔다. FlashAttention은 표에서 가장 나쁜 복잡도를 갖고 있지만 실제로 가장 널리 쓰인다. 상수가 크게 다르고, 무엇보다 품질을 내주지 않기 때문이다.
기법의 조합
이들이 서로 배타적이라고 오해하기 쉬운데 그렇지 않다. 희소 패턴은 어떤 셀을 계산할지를 정하고 FlashAttention은 그 셀들을 어떻게 옮길지를 정하므로, 둘은 다른 층위에 있다. 슬라이딩 윈도를 쓰는 모델도 그 윈도 안의 어텐션은 타일링 커널로 돈다. 실제로 널리 쓰이는 커널들이 윈도 옵션을 인자로 받는 이유가 이것이다.
SSM 쪽도 마찬가지다. Mamba의 선택적 재귀를 실용적인 속도로 돌게 만든 것은 병렬 스캔을 SRAM 안에서 끝내는 커널이었고, 이는 FlashAttention과 같은 원리를 다른 연산에 적용한 것이다. 효율화의 진짜 교훈은 특정 패턴이 아니라 메모리 계층을 의식하고 알고리즘을 짜는 태도에 있다.
KV 캐시
이 글에서 줄인 것은 어텐션을 계산하는 비용이다. 학습에서는 이것이 거의 전부지만, 서비스 중인 모델의 추론에서는 다른 항이 하나 더 크게 남는다. 토큰을 하나씩 만들어 내는 동안 이전 토큰들의 키와 값을 계속 들고 있어야 하고, 이 저장소는 배치와 컨텍스트 길이에 비례해 자란다. 앞에서 슬라이딩 윈도의 이점으로 잠깐 언급한 바로 그 캐시다.
FlashAttention은 여기에 손대지 않는다. 계산 중 만들어지는 중간 행렬을 없앨 뿐, 다음 스텝을 위해 남겨 둬야 하는 값은 그대로다. 다음 글에서는 이 저장소를 직접 줄이는 접근을 본다 — 헤드마다 따로 두던 키와 값을 여러 헤드가 나눠 쓰게 만들어, 품질을 거의 잃지 않으면서 캐시를 몇 분의 일로 접는 방법이다.
읽어주셔서 감사합니다. 😊

