FlashAttention 논문 (IO-awareness·타일링)
FlashAttention 논문 (IO-awareness·타일링)
FlashAttention 논문(2022, Tri Dao 외)은 "왜 어텐션이 느린가"를 딱 짚어요 — Transformer가 긴 시퀀스에서 느리고 메모리를 많이 먹는 이유는 어텐션의 시간·메모리 복잡도가 **시퀀스 길이의 제곱(O(n²))**이기 때문이에요.
근사 어텐션 방법은 모델 품질을 희생해 연산 복잡도를 줄이려 했지만 벽시계 속도(wall-clock) 개선을 잘 못 냈어요. 저자들은 빠진 원칙이 IO-awareness라고 주장하며, GPU 메모리 계층(HBM ↔ SRAM) 사이의 읽기·쓰기를 인지하는 정확 어텐션 알고리즘을 제안해요.
타일링으로 HBM 접근 줄이기
FlashAttention은 tiling을 써서 GPU HBM과 온칩 SRAM 사이의 메모리 읽기·쓰기 횟수를 줄여요.
- 표준 어텐션: 소프트맥스 계산을 위해 어텐션 점수 행렬 전체를 메모리에 둠 → HBM 접근이 많음.
- FlashAttention: 블록으로 나눠(tiling) SRAM에 유지하면서 소프트맥스를 온라인 스트리밍 방식으로 갱신 → HBM 접근 최소화.
- IO 복잡도 분석 결과, 표준 어텐션보다 HBM 접근이 적고 여러 SRAM 크기에서 최적.
실측 속도 개선 (논문)
- BERT-large(시퀀스 512): MLPerf 1.1 기록 대비 15% end-to-end 속도 향상.
- GPT-2(시퀀스 1K): 3× 속도향상.
- Long Range Arena(시퀀스 1K~4K): 2.4×.
- Path-X(시퀀스 16K)에서 첫 chance 이상 성능 달성 — 새 기능.
블록 스파스 어텐션 확장
FlashAttention은 block-sparse 어텐션으로도 확장돼, 기존 근사 어텐션보다 빠른 근사 알고리즘을 제공해요.