FlashAttention 논문 (IO-awareness·타일링)

FlashAttention 논문 (IO-awareness·타일링)

FlashAttention 논문(2022, Tri Dao 외)은 "왜 어텐션이 느린가"를 딱 짚어요 — Transformer가 긴 시퀀스에서 느리고 메모리를 많이 먹는 이유는 어텐션의 시간·메모리 복잡도가 **시퀀스 길이의 제곱(O(n²))**이기 때문이에요.

근사 어텐션 방법은 모델 품질을 희생해 연산 복잡도를 줄이려 했지만 벽시계 속도(wall-clock) 개선을 잘 못 냈어요. 저자들은 빠진 원칙이 IO-awareness라고 주장하며, GPU 메모리 계층(HBM ↔ SRAM) 사이의 읽기·쓰기를 인지하는 정확 어텐션 알고리즘을 제안해요.

출처: https://arxiv.org/abs/2205.14135

타일링으로 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 어텐션으로도 확장돼, 기존 근사 어텐션보다 빠른 근사 알고리즘을 제공해요.

더 알아보기