본문 바로가기
Inference/Optimization

FlashAttention

by AteN 2024. 8. 27.

FlashAttention 완전 정리, 근사 없이 같은 값을 더 빠르게

요약

  • FlashAttention은 어텐션을 수학적으로 똑같이 계산하되, 큰 n x n 점수 행렬을 느린 GPU 메모리에 만들지 않고 빠른 on-chip 메모리에서 타일 단위로 처리하는 알고리즘입니다.
  • 가장 자주 헷갈리는 지점. 근사가 아닙니다. 스파스나 선형 어텐션은 품질을 일부 포기해 속도를 얻지만, FlashAttention은 표준 어텐션과 같은 값을 냅니다.
  • 출발점이 된 관찰이 핵심입니다. 병목은 연산이 아니라 메모리 이동입니다. 논문은 이를 "IO-aware"라 부르고, 그 이동을 줄이도록 알고리즘을 설계했습니다.
  • 수학적 열쇠는 online softmax입니다. softmax는 원래 행 전체가 필요한데, 부분합을 보정하며 누적해 전체 행렬 없이 정확한 결과를 만듭니다.
  • 버전마다 개선 지점이 다릅니다. FA-1은 IO, FA-2는 작업 분할, FA-3는 Hopper의 비동기와 FP8, FA-4는 CuTeDSL로 Hopper와 Blackwell을 겨냥합니다.
  • 실무 함정. GPU 아키텍처와 dtype에 의존하고, 조용히 표준 어텐션으로 폴백되는 경우가 있습니다.

1장. 왜 필요했나

1.1 문제 설정

표준 어텐션은 n x n 점수 행렬을 만들어 GPU 메모리에 쓰고 다시 읽습니다. 여기서 두 문제가 생깁니다.

① 느리다
   연산이 아니라 메모리 대역폭에 묶여 있다 (memory-bound)
   GPU 가 계산할 여력은 남는데 데이터를 옮기느라 놀고 있다

② 메모리를 폭식한다
   n 제곱 행렬이 길이의 제곱으로 커진다
   긴 컨텍스트가 아예 불가능해진다

FlashAttention은 이 n x n 행렬을 애초에 만들지 않습니다.

1.2 기존 접근과 무엇이 다른가

논문 서문이 이 대비를 명확히 합니다.

근사 어텐션 (스파스, 선형 등)
   연산 복잡도를 줄이려고 모델 품질을 맞바꿨다
   그런데 실제 벽시계 시간(wall-clock) 가속은 자주 얻지 못했다

   왜냐하면 FLOPs 를 줄여도
   병목이 FLOPs 가 아니었기 때문이다

FlashAttention 의 주장
   빠진 원칙은 어텐션 알고리즘을 IO-aware 하게 만드는 것이다
   즉 GPU 메모리 계층 간 읽기와 쓰기를 계산에 넣는 것

"FLOPs를 줄였는데 왜 안 빨라지나"에 대한 답이 이 논문의 출발점입니다. 그래서 FlashAttention은 연산량을 줄이지 않고도 빨라집니다.

1.3 흔한 오해

❌ "FlashAttention 은 희소 어텐션이다"
❌ "정확도를 조금 포기하고 속도를 얻는 기법이다"

   둘 다 아니다. exact 하다
   표준 어텐션과 동일한 값을 낸다

   다만 논문은 별도로 block-sparse FlashAttention 도 제시하는데
   이건 명시적으로 근사 알고리즘이다
   둘을 구분해야 한다

2장. GPU 메모리 계층

2.1 왜 계층을 알아야 하나

FlashAttention을 이해하려면 GPU 메모리가 비대칭 계층이라는 걸 먼저 알아야 합니다.

HBM (High Bandwidth Memory)
   흔히 말하는 "GPU 메모리". 수십 GB
   크지만 상대적으로 느리다

SRAM (on-chip shared memory)
   SM 안에 있는 작업 공간. 수십에서 수백 KB 급
   HBM 보다 훨씬 빠르지만 매우 작다

   용량과 속도가 반대로 간다
   FlashAttention 은 이 비대칭을 이용한다

2.2 작은 작업대 비유

거대한 표를 한 번에 펼칠 작업대가 없다

   표를 작은 타일로 나눠
   작업대에 하나씩 올려 처리하고
   부분 결과를 누적한다

   큰 표를 창고에 보관할 필요가 없다

   작업대 = SRAM (빠르고 작다)
   창고   = HBM (크고 느리다)

2.3 IO 복잡도라는 관점

논문은 FlashAttention의 IO 복잡도를 분석해, 표준 어텐션보다 HBM 접근이 적고 일정 범위의 SRAM 크기에서는 최적임을 보였습니다.

"몇 번 계산하나"가 아니라 "몇 번 읽고 쓰나"로 알고리즘을 평가한다는 관점 자체가 이 논문의 기여입니다.


3장. 어떻게 동작하나

3.1 전체 흐름

1. Q, K, V 를 타일(블록)로 나눠 SRAM 에 올린다
   │
   ▼
2. 타일끼리 점수와 softmax 와 가중합을 계산한다
   단 online softmax 로 부분합을 정규화하며 누적한다
   │
   ▼
3. 전체 n x n 행렬을 HBM 에 만들지 않는다
   │
   └──► HBM 왕복이 급감한다

3.2 online softmax, 진짜 어려운 부분

2번이 이 알고리즘의 수학적 핵심입니다.

3.2.1 왜 어려운가

softmax 는 보통 행 전체가 필요하다

   수치 안정성을 위해 최댓값을 빼고
   지수를 취한 뒤
   합으로 나눈다

   전체를 안 보면 최댓값도 합도 모른다
   그런데 타일 단위로 처리하면 전체를 못 본다

3.2.2 해법

타일을 순회하며
running max 와 running sum 으로 부분 softmax 를 점진 보정한다

   새 타일에서 더 큰 값이 나오면
   이전까지의 결과를 그 새 최댓값에 맞춰 재조정한다

   구체적으로는
   이전 합과 이전 출력에 보정 계수를 곱해 스케일을 맞춘 뒤
   새 타일의 기여를 더한다

결과적으로 전체 행렬 없이도 정확한 값이 나옵니다. 근사가 아닌 이유가 여기 있습니다.

이 기법 자체는 FlashAttention이 처음 만든 게 아닙니다. 메모리 접근을 줄여 softmax를 계산하는 online softmax가 앞선 연구로 있었고, FlashAttention이 이를 어텐션 전체로 확장했습니다.

3.3 역전파에서의 재계산

학습할 때는 역전파에서 큰 중간값이 필요합니다.

표준 구현   forward 에서 만든 n x n 행렬을 저장해 뒀다가 backward 에서 쓴다
            메모리가 n 제곱으로 든다

FlashAttention   저장하지 않고 backward 때 다시 계산한다
                 메모리를 선형급으로 유지한다

   약간의 추가 연산과 메모리를 맞바꾼 것이다
   그런데 병목이 메모리 이동이라 전체적으로는 이득이다

"다시 계산하는 게 저장했다 읽는 것보다 싸다"는 게 직관에 반하지만, 메모리 대역폭이 병목인 환경에서는 맞습니다.

3.4 세 가지 아이디어로 정리

① IO-awareness
   FLOPs 가 아니라 HBM 과 SRAM 사이 이동을 줄이도록 설계
   연산량은 비슷하거나 약간 늘어도 실측이 빨라진다

② 타일링 + online softmax
   전체 행렬 없이 정확한 softmax 를 만드는 수학적 장치

③ backward 재계산
   중간값을 저장하는 대신 다시 계산해 메모리를 선형으로

4장. 버전별로 무엇이 달라졌나

4.1 FlashAttention (2022)

IO-aware exact attention을 제시했습니다. 논문이 보고한 수치입니다.

학습 가속
   BERT-large (길이 512)   MLPerf 1.1 기록 대비 엔드투엔드 15% 단축
   GPT-2 (길이 1K)         3배
   Long-range arena (1K~4K) 2.4배

긴 컨텍스트가 열어 준 것
   GPT-2 에서 perplexity 0.7 개선
   장문 분류에서 6.4 포인트 향상

   Path-X (길이 16K)    61.4% 정확도
   Path-256 (길이 64K)  63.1% 정확도
   → 이 두 과제에서 우연 이상 성능을 낸 최초의 트랜스포머

마지막 항목이 인상적입니다. 속도 개선이 아니라 이전에는 불가능하던 과제를 가능하게 만들었습니다.

4.2 FlashAttention-2 (2023)

FA-1의 한계를 정확히 지목하며 출발합니다.

FA-1 은 여전히 최적화된 행렬곱(GEMM)만큼 빠르지 않았다
   이론상 최대 FLOPs/s 의 25~40% 에 그쳤다

원인은 알고리즘이 아니라 GPU 상의 작업 분할이었다
   thread block 과 warp 사이 분할이 최적이 아니라
   낮은 점유율(occupancy) 이나
   불필요한 shared memory 읽기와 쓰기를 유발했다

해법 셋입니다.

① 비matmul FLOPs 를 줄이도록 알고리즘을 손봤다
   Tensor Core 는 matmul 에 특화돼 있어
   그 외 연산이 상대적으로 훨씬 비싸다

② 단일 헤드에서도 thread block 을 가로질러 병렬화해
   점유율을 높였다

③ thread block 안에서 warp 사이 작업을 재분배해
   shared memory 통신을 줄였다

결과입니다.

FA-1 대비 약 2배
A100 에서 이론상 최대 FLOPs/s 의 50~73% 도달
GPT 계열 학습에서 A100 당 최대 225 TFLOPs/s (모델 FLOPs 활용률 72%)

FA-2의 기여는 알고리즘이 아니라 GPU 매핑이라는 점이 중요합니다. 같은 수학인데 하드웨어에 어떻게 얹느냐로 2배가 갈렸습니다.

4.3 FlashAttention-3 (2024)

같은 패턴이 반복됩니다. 새 하드웨어가 나오자 다시 활용률이 떨어졌습니다.

FA-2 는 H100 에서 35% 활용률에 그쳤다
   Hopper 의 새 기능을 쓰지 않았기 때문이다

Hopper용 기법 셋입니다.

① warp-specialization 으로 연산과 데이터 이동을 겹친다
   Tensor Core 와 TMA 의 비동기성을 활용

② 블록 단위 matmul 과 softmax 를 교차 배치한다
   서로 다른 유닛이 동시에 일하게 한다

③ 블록 양자화와 incoherent processing 으로
   FP8 저정밀 하드웨어 지원을 활용한다

결과입니다.

H100 에서 FP16 으로 1.5~2.0배 가속, 최대 740 TFLOPs/s (75% 활용률)
FP8 로는 1.2 PFLOPs/s 에 근접

FP8 FlashAttention-3 는 기준 FP8 어텐션 대비
수치 오차가 2.6배 낮다

마지막 줄이 실무적으로 중요합니다. FP8은 정밀도가 낮아 오차가 걱정되는데, 블록 양자화와 incoherent processing이 그 오차를 줄였습니다. "저정밀을 쓰되 오차는 관리한다"는 접근입니다.

4.4 FlashAttention-4

공식 저장소 기준으로 CuTeDSL로 작성되었고 Hopper와 Blackwell을 겨냥합니다.

pip install flash-attn-4

from flash_attn.cute import flash_attn_func
out = flash_attn_func(q, k, v, causal=True)

CuTeDSL로 옮겨간 게 흐름상 의미가 있습니다. 세대마다 손으로 CUDA를 다시 튜닝하던 것을, 더 높은 수준의 DSL로 기술해 하드웨어 세대 전환 비용을 줄이려는 방향입니다.

4.5 관통하는 패턴

FA-1   알고리즘 차원의 IO 최적화
FA-2   GPU 작업 분할 최적화
FA-3   특정 세대(Hopper) 하드웨어 기능 활용
FA-4   DSL 로 기술해 세대 전환 비용을 낮춤

   버전이 올라갈수록 하드웨어에 더 밀착한다
   그래서 최신 버전일수록 요구 GPU 가 좁아진다

이게 다음 장의 실무 주의사항으로 이어집니다.


5장. 실무

5.1 어느 버전이 내 GPU에서 되나

공식 저장소 기준입니다.

FlashAttention-2 (CUDA)
   Ampere, Ada, Hopper 계열   예: A100, RTX 3090, RTX 4090, H100
   Turing (T4, RTX 2080) 은 별도 저장소에서 일부 기능만 지원
   dtype 은 fp16 과 bf16. bf16 은 Ampere 이상 필요
   head dimension 은 256 까지

FlashAttention-2 (ROCm)
   composable kernel 백엔드와 Triton 백엔드 두 가지

FlashAttention-3
   H100 또는 H800 필요
   FP16 과 BF16 은 forward 와 backward, FP8 은 forward

FlashAttention-4
   Hopper 와 Blackwell 겨냥

"FlashAttention을 켰는데 안 빨라진다"의 첫 번째 원인이 아키텍처 불일치입니다.

5.2 켜는 법

5.2.1 transformers

from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained(
    model_id,
    attn_implementation="flash_attention_2",
    dtype="bfloat16",          # fp32 로는 안 된다
)

5.2.2 서빙 엔진

vLLM과 TGI는 FlashAttention 계열 커널을 내장하고 있어 별도로 켤 필요가 없습니다. MLA용 전용 커널처럼 어텐션 변형별 커널도 따로 나와 있습니다.

5.2.3 llama.cpp

llama-server -m model.gguf -fa on

여기에 연쇄 효과가 하나 있습니다. V 캐시를 양자화하려면 flash attention이 켜져 있어야 합니다. -fa 없이 -ctv q4_0을 주면 컨텍스트 생성 자체가 실패합니다.

5.3 확인할 것 셋

GPU 아키텍처   5.1 의 표와 대조
dtype          fp16 이나 bf16. fp32 는 지원 안 되는 경우가 있다
마스크 조합     일부 커스텀 마스크나 어텐션 바이어스는 미지원

설치했는데 조용히 표준 어텐션으로 폴백되는 경우가 있습니다. 로그를 보거나 메모리 사용량이 실제로 줄었는지로 확인하는 게 안전합니다.

5.4 prefill과 decode 중 어디에 효과가 있나

실무에서 기대와 결과가 갈리는 지점입니다.

FlashAttention은 어텐션 계산의 메모리 이동을 줄입니다. 그런데 어텐션이 전체 비용에서 차지하는 비중은 단계마다 다릅니다.

prefill
   토큰이 많아 어텐션이 n 제곱으로 커진다
   어텐션이 전체 비용을 지배한다
   └──► 효과가 크다

decode
   토큰이 1개라 어텐션 비중이 작다
   가중치를 HBM 에서 읽는 게 지배적이다
   └──► 효과가 작다

그래서 긴 프롬프트를 처리할 때와 학습할 때 이득이 큽니다. 짧은 프롬프트에 긴 생성이면 체감이 덜합니다.

컨텍스트가 길어질수록 효과가 커진다고 보면 대체로 맞습니다. 4.1의 Path-X 같은 초장문 과제에서 가장 극적이었던 이유도 같습니다.

5.5 왜 거의 모든 서빙이 쓰나

같은 모델을 더 적은 메모리로 더 빠르게 돌린다
품질 손실이 없다
바꿀 게 커널 하나뿐이다

   메모리가 선형이 되니 컨텍스트를 늘릴 여지가 생기고
   그 여지가 KV 캐시 예산으로 돌아간다

긴 컨텍스트 학습과 추론에서는 사실상 필수가 됐습니다.


6장. 정리

6.1 핵심 셋

① n x n 행렬을 만들지 않고 타일로 계산하는 정확하고 빠른 어텐션이다
   근사가 아니다

② 병목이 연산이 아니라 메모리 이동이라는 관찰에서 출발했다
   연산량이 비슷해도 실측이 빨라진다

③ 버전마다 개선 지점이 다르다
   FA-1 알고리즘, FA-2 작업 분할, FA-3 Hopper 기능, FA-4 DSL
   최신일수록 요구 GPU 가 좁아진다

6.2 남는 질문

  • 하드웨어 세대마다 커널을 다시 짜는 비용을 어디까지 줄일 수 있나. FA-4의 DSL 접근이 그 방향이지만 아직 초기입니다.
  • FP8과 그 아래 정밀도에서 오차를 어디까지 관리할 수 있나. FA-3가 2.6배 개선을 보였지만 정밀도를 더 낮추면 다시 문제가 됩니다.
  • decode 단계의 어텐션도 개선 여지가 있나. 지금 이득은 prefill에 쏠려 있는데, decode의 병목은 성격이 다릅니다.

6.3 용어 정리

용어 한 줄 뜻
FlashAttention n x n 행렬을 안 만들고 타일로 계산하는 정확하고 빠른 어텐션
exact 근사가 아니라 표준 어텐션과 동일한 결과
IO-aware 연산이 아니라 메모리 계층 간 읽기와 쓰기를 줄이도록 설계하는 접근
HBM / SRAM 느리고 큰 GPU 메모리 / 빠르고 작은 on-chip 메모리
memory-bound 연산이 아니라 메모리 대역폭에 성능이 묶인 상태
타일링 큰 행렬을 작은 블록으로 나눠 처리하는 것
online softmax 부분합을 정규화하며 누적해 전체 행렬 없이 softmax를 계산
running max / sum 타일을 순회하며 유지하는 최댓값과 합. 보정의 기준
재계산 (recomputation) 역전파에서 중간값을 저장하는 대신 다시 계산해 메모리 절약
occupancy GPU 연산 유닛이 얼마나 채워져 일하는지의 척도
warp-specialization warp마다 다른 역할을 맡겨 연산과 데이터 이동을 겹치는 기법
TMA Hopper의 비동기 메모리 전송 유닛
block-sparse FlashAttention 원 논문이 함께 제시한 별도의 근사 알고리즘

6.4 참고자료

'Inference > Optimization' 카테고리의 다른 글

KV Cache 란  (0) 2024.07.02
prefill vs decode  (0) 2024.03.11

댓글