FlashAttention 라이브러리 개요

FlashAttention 라이브러리 개요

FlashAttention은 빠르고 메모리 효율적인 정확 어텐션 알고리즘의 공식 구현 저장소예요. FlashAttention과 FlashAttention-2 논문의 코드를 제공하고, 이후 3/4 버전을 베타로 내놓고 있어요.

기본 사용법은 아주 간단해서, Q·K·V 텐서만 넘기면 돼요. flash_attn_func(q, k, v, causal=True)처럼 호출하면 됩니다. 주의할 점은 CUDA/ROCm 툴킷과 PyTorch 2.2+가 필요하고, 실제 속도는 GPU 아키텍처에 크게 의존해요.

출처: https://github.com/Dao-AILab/flash-attention

설치

pip install flash-attn --no-build-isolation

메모리가 적은 머신에서는 병렬 컴파일을 제한할 수 있어요.

MAX_JOBS=4 pip install flash-attn --no-build-isolation

기본 인터페이스

from flash_attn import flash_attn_qkvpacked_func, flash_attn_func

flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=False,
                window_size=(-1, -1), alibi_slopes=None, deterministic=False)
  • causal — 자기회귀 모델용 인과 마스크.
  • window_size — 슬라이딩 윈도우 로컬 어텐션(i 위치는 [i-left, i+right] 키만 참조).
  • MQA/GQA — K·V의 헤드 수를 Q보다 적게 넘기면 그룹 어텐션 지원.

성능·메모리

  • 메모리 절감은 시퀀스 길이에 비례 — 표준 어텐션은 메모리가 O(n²), FlashAttention은 O(n).
  • 시퀀스 2K에서 약 10배, 4K에서 약 20배 메모리 절감 → 훨씬 긴 시퀀스 확장 가능.
  • H100 기준 FP16/BF16에서 고속 스루풋.

더 알아보기