FlexAttention 소개 — 유연성과 성능
FlexAttention 소개 — 유연성과 성능
PyTorch 블로그에서 처음 소개된 FlexAttention은 "FlashAttention의 성능 + PyTorch의 유연성"이라는 목표를 가진 API예요. 기존에는 새 어텐션 변형을 시도할 때마다 커스텀 CUDA 커널을 직접 짜야 했지만, FlexAttention은 이 문제를 풀었어요.
배경: '어텐션 변형의 하이퍼큐브 문제'
- 연구자들은 ALiBi, 접두사 LM, 슬라이딩 윈도우, 문서 마스킹 등 많은 어텐션 변형을 시도해요.
- FlashAttention류 융합 커널은 빠르지만 패턴이 고정되어 있어서, 변형마다 새 커널을 작성해야 했어요.
- FlexAttention은
score_mod라는 사용자 함수 하나로 이러한 변형을 표현하게 만듭니다.
핵심 아이디어
score_mod(score, b, h, q_idx, kv_idx)같은 함수로 어텐션 점수를 수정해요.torch.compile이 이 함수를 단일 융합 FlexAttention 커널로 낮춰서, 불필요한 메모리 생성 없이 실행돼요.- 기존 손으로 만든 커널과 견줄 만한 성능을 보여줍니다.
사용 예감
from torch.nn.attention.flex_attention import flex_attention
def noop(score, b, h, q, kv): return score
out = flex_attention(q, k, v, score_mod=noop)