Zettelkasten

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의 두 가지 아이디어

  1. 타일링: Q, K, V를 작은 블록으로 쪼개 칩 안의 작고 빠른 메모리(SRAM)에 올리고, 점수 계산 → softmax → V 곱까지 커널 하나 안에서 끝낸다. N×N 행렬을 HBM에 쓰지 않는다
  2. 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 등으로 증명할 수 있으면, 그 바깥의 방어적 복사는 줄여도 안전할 가능성이 높다 — 다른 커널에도 일반화되는지는 아직 검증 안 됨

참고