[Daily morning study] Flash Attention 원리와 IO-Aware 어텐션 최적화
#daily morning study
문제: Attention은 왜 느린가?
Transformer의 셀프 어텐션은 시퀀스 길이 N에 대해 O(N²)의 시간 복잡도와 메모리를 요구한다. 시퀀스가 길어질수록 계산량과 메모리가 폭발적으로 증가해서 긴 컨텍스트를 처리하는 데 병목이 생긴다.
기존 어텐션 연산을 수식으로 보면:
Attention(Q, K, V) = softmax(QKᵀ / √d_k) · V
Q, K, V는 각각 (N, d_k) 크기의 행렬이고, QKᵀ는 (N, N) 행렬이 된다. N=8192인 경우 8192×8192 = 약 67M 요소를 저장해야 한다.
GPU 메모리 계층 구조
Flash Attention을 이해하려면 GPU 메모리 구조부터 알아야 한다.
| 메모리 종류 | 크기 | 대역폭 | 특징 |
|---|---|---|---|
| SRAM (온칩) | ~수십 MB | ~19 TB/s | 빠르지만 작음 |
| HBM (글로벌 메모리) | ~40-80 GB | ~2 TB/s | 크지만 느림 |
연산 유닛(SM, Streaming Multiprocessor)은 SRAM에서 데이터를 읽어 연산하고, 결과를 다시 HBM에 쓴다. SRAM ↔ HBM 간 데이터 이동(메모리 I/O)이 병목이다.
기존 어텐션은:
- QKᵀ 행렬을 HBM에 쓴다 (N×N 크기)
- softmax를 위해 다시 HBM에서 읽는다
- 결과에 V를 곱해 HBM에 쓴다
- dropout, masking마다 추가 HBM 접근
총 HBM 접근 횟수: O(N²) — 여기서 실제 속도 저하가 발생한다.
Flash Attention의 핵심 아이디어
Flash Attention(2022, Tri Dao et al.)은 타일링(Tiling) 기법으로 이 문제를 해결한다.
핵심 원리: QKᵀV 연산 전체를 SRAM 안에서 작은 블록으로 나눠 처리하고, N×N 행렬을 HBM에 쓰지 않는다.
온라인 softmax (Online Softmax)
softmax는 전체 벡터의 최댓값과 합계를 알아야 정규화할 수 있다:
softmax(xᵢ) = exp(xᵢ) / Σ exp(xⱼ)
블록별로 처리하려면 전체를 보지 않고 softmax를 계산해야 한다. 이를 위해 수치 안정성을 유지하면서 블록을 순회하며 점진적으로 softmax를 업데이트하는 online softmax 알고리즘을 사용한다.
각 블록마다 현재까지의 최댓값(m)과 정규화 인수(l)를 추적하며 accumulate한다:
m_new = max(m_old, max(x_block))
l_new = exp(m_old - m_new) * l_old + sum(exp(x_block - m_new))
O_new = diag(exp(m_old - m_new)) * O_old + exp(x_block - m_new) * V_block
최종적으로 O / l이 정확한 softmax 결과와 일치한다.
타일링 알고리즘
for i in range(0, N, BLOCK_SIZE): # Q 블록 반복
Qᵢ = Q[i:i+BLOCK_SIZE] # SRAM 로드
m, l, O = -inf, 0, 0
for j in range(0, N, BLOCK_SIZE): # K, V 블록 반복
Kⱼ = K[j:j+BLOCK_SIZE] # SRAM 로드
Vⱼ = V[j:j+BLOCK_SIZE] # SRAM 로드
S = Qᵢ · Kⱼᵀ / √d # 내적 계산 (SRAM 내)
m, l, O = update(m, l, O, S, Vⱼ) # 온라인 업데이트
O = O / l # 최종 정규화
HBM에 O만 저장 # 한 번만 씀
HBM 접근 횟수가 O(N²)에서 O(N)으로 줄어든다.
Flash Attention vs 기존 어텐션 비교
| 항목 | 기존 어텐션 | Flash Attention |
|---|---|---|
| HBM 접근 | O(N²) | O(N) |
| 메모리 사용 | O(N²) | O(N) |
| 계산 복잡도 | O(N²d) | O(N²d) (동일) |
| 실제 속도 | 기준 | 2~4배 빠름 |
| 역전파 가능 | 가능 | 가능 (재계산 방식) |
계산량(FLOPs)은 동일하지만, 메모리 I/O가 줄어서 실제 속도가 크게 빨라진다. 이런 관점에서 Flash Attention을 IO-Aware 어텐션 알고리즘이라 부른다.
역전파에서의 재계산 (Recomputation)
역전파 시 gradient를 계산하려면 순전파의 어텐션 행렬 S와 P(= softmax(S))가 필요하다. 기존에는 이것을 저장해뒀지만, Flash Attention은 저장하지 않는다.
대신 역전파 시에 순전파 계산을 재계산(recomputation)한다. 이렇게 하면:
- 메모리 절약: O(N²) → O(N)
- 추가 연산: FLOP 수 약 1.5배 증가
- 결과적으로 메모리 절약 이득이 연산 증가보다 크다
Flash Attention 2와 3
Flash Attention 2 (2023)
- 작업 분배 최적화: 워프(warp) 간 통신 감소
- 내부 루프 순서 변경으로 비공유 메모리 접근 최소화
- Flash Attention 대비 2배 추가 속도 향상
- 시퀀스 병렬성(sequence parallelism) 지원
Flash Attention 3 (2024)
- Hopper GPU(H100)의 새 기능 활용 (Tensor Memory Accelerator, WGMMA)
- 비동기 파이프라이닝으로 메모리 복사와 연산을 겹침
- FP8 지원으로 추가 속도 향상
- Flash Attention 2 대비 약 1.5~2배 향상
Multi-Query Attention (MQA)과 Grouped Query Attention (GQA)와의 관계
Flash Attention은 어텐션 연산의 메모리 I/O를 줄이는 반면, MQA/GQA는 KV 헤드 수를 줄여 KV 캐시 크기 자체를 줄이는 다른 접근이다.
| 기법 | Q 헤드 | K/V 헤드 | 특징 |
|---|---|---|---|
| Multi-Head Attention (MHA) | H | H | 품질 최고 |
| Multi-Query Attention (MQA) | H | 1 | 속도 최고, 품질 약간 저하 |
| Grouped Query Attention (GQA) | H | G (1 < G < H) | 균형 |
Llama 2/3, Mistral 등 최신 모델은 GQA를 기본으로 사용한다. Flash Attention과 GQA는 서로 보완적으로 함께 사용할 수 있다.
실제 사용 예시
# PyTorch 2.0+ - scaled_dot_product_attention에 Flash Attention 통합
import torch
import torch.nn.functional as F
# Flash Attention이 자동으로 활용됨 (CUDA 환경)
output = F.scaled_dot_product_attention(
query, key, value,
attn_mask=None,
dropout_p=0.0,
is_causal=True # causal masking
)
# 명시적으로 Flash Attention 라이브러리 사용
from flash_attn import flash_attn_qkvpacked_func
qkv = torch.randn(batch_size, seq_len, 3, num_heads, head_dim, device='cuda', dtype=torch.float16)
output = flash_attn_qkvpacked_func(qkv, dropout_p=0.0, causal=True)
PyTorch 2.0부터 scaled_dot_product_attention이 내부적으로 Flash Attention을 자동으로 선택한다.
정리
- Flash Attention은 어텐션 연산의 HBM I/O를 O(N²)에서 O(N)으로 줄임
- 타일링 + 온라인 softmax로 중간 N×N 행렬을 HBM에 쓰지 않음
- 계산량은 동일하지만 실제 속도는 2~4배 향상, 메모리는 O(N²)→O(N)
- Flash Attention 2, 3로 진화하며 지속적으로 개선 중
- LLM 학습과 추론 모두에서 긴 컨텍스트 처리의 핵심 기술