[Daily morning study] Flash Attention 원리와 IO-Aware 어텐션 최적화

#daily morning study

Image


문제: 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)이 병목이다.

기존 어텐션은:

  1. QKᵀ 행렬을 HBM에 쓴다 (N×N 크기)
  2. softmax를 위해 다시 HBM에서 읽는다
  3. 결과에 V를 곱해 HBM에 쓴다
  4. 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)HH품질 최고
Multi-Query Attention (MQA)H1속도 최고, 품질 약간 저하
Grouped Query Attention (GQA)HG (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 학습과 추론 모두에서 긴 컨텍스트 처리의 핵심 기술