Flash Attention은 어텐션을 칩 안의 메모리에서 블록 단위로 계산해 큰 메모리 왕복을 없앤다
·수정 1회
요약
- Flash Attention은 큰 연산에서 메모리를 왕복하면서 생기는 비효율을, 작은 연산으로 쪼개 칩 안(SRAM)에서 끝내는 방식으로 없애 더 빠르게 만든다
- softmax는 블록마다 "지금까지의 최댓값"을 들고 다니며 계산하기 때문에, 쪼개도 결과가 근사가 아니라 정확히 같다
본문
일반 어텐션이 느린 이유
- 계산은
S = Q·Kᵀ→P = softmax(S)→O = P·V - 길이 N이면 중간 행렬 S, P가 N×N. 일반 구현은 이걸 GPU의 큰 메모리(HBM)에 썼다가 다시 읽는다
- 병목은 계산량이 아니라 메모리 왕복이다. 메모리도 N²으로 늘어난다
Flash Attention의 두 가지 아이디어
- 타일링: Q, K, V를 작은 블록으로 쪼개 칩 안의 작고 빠른 메모리(SRAM)에 올리고, 점수 계산 → softmax → V 곱까지 커널 하나 안에서 끝낸다. N×N 행렬을 HBM에 쓰지 않는다
- online softmax: softmax는 원래 한 행 전체의 최댓값과 합이 필요하다. 블록마다 세 값을 들고 다닌다
- 지금까지의 최댓값
m, 지수 합l, 출력 누적O - 새 블록에서 더 큰 최댓값
m'이 나오면 기존l,O에e^(m − m')을 곱해 새 기준으로 다시 스케일링한 뒤 새 블록 값을 더한다 - 마지막에
O / l→ 전체를 한 번에 본 softmax와 정확히 같다 - 비유: 환율이 바뀔 때마다 모은 돈을 새 환율로 다시 환산해 두는 것
- 지금까지의 최댓값
- 결과: 메모리 O(N²) → O(N), 메모리 왕복이 줄어 긴 시퀀스일수록 빨라진다
CTranslate2 Whisper 포크 작업에서 만난 지점
- 업스트림 인코더 버그: v4.8.2는
flash_attention=True를 인코더에 전달하지 않아, flash-on 벤치 48구성 중 22구성에서 업스트림 토큰이 달랐다 (포크 라운드 1에서 수정) - 디코더 KV 캐시와
seqlens_k: flash 경로는 self-attention KV 캐시를 515칸(3+512)으로 미리 잡고, 커널은seqlens_k= 유효 길이 L까지만 읽는다. 그런데 beam 재정렬(gather)은 매 스텝 515칸 전체를 복사하고 있었다 — whisper-small beam5 기준 디코드 GPU 시간 약 50ms 중 5.9ms - 라운드 6 prefix gather: 커널이 L 뒤를 읽지 않는 걸 sanitizer로 확인한 뒤 유효 L칸만 복사하게 바꿨다 → gather 370us → 7.3us/launch, flash beam5 전체 -9~-12% (A10G, 토큰 동일)
- flash가 아닌 경로와의 차이: graphs/패딩 모드의 일반 어텐션은 캐시 뒤 빈칸도 마스크를 씌워 읽는다. 그래서 그냥 잘라낼 수 없고, 빈칸이 항상 0으로 유지되도록 보장하는
keep_tail방식이 필요했다 - 교훈(이 사례에서의 추정): 커널이 실제로 어디까지 읽는지를 sanitizer 등으로 증명할 수 있으면, 그 바깥의 방어적 복사는 줄여도 안전할 가능성이 높다 — 다른 커널에도 일반화되는지는 아직 검증 안 됨
참고
- Transformer에서 self attention 계산 과정 — Flash Attention이 최적화하는 QKᵀ → softmax → PV 계산 자체
- Encoder에서 self attention 계산시 Lower Triangular를 이용해 인과적 구조로 예측가능한 트랜스포머 구조를 만든다. — 인과적 구조라 이전 KV를 재사용 → 디코더 KV 캐시의 전제
- ctranslate2 num_workers는 가중치를 공유하고 worker별 메모리는 호출 중에만 늘어난다 — 같은 CT2 디코더의 KV 캐시·beam state 메모리
- ctranslate2의 GIL 해제와 num_workers는 독립적인 두 메커니즘이다 — 같은 CT2/faster-whisper 실험 시리즈
- Whisper 음성 처리와 최적화 방식 — faster-whisper(CTranslate2) 가속 원리 배경