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 참고자료
- Dao et al., "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness" (NeurIPS 2022, arXiv:2205.14135)
- Dao, "FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning" (arXiv:2307.08691)
- Shah et al., "FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision" (arXiv:2407.08608)
- Milakov, Gimelshein, "Online normalizer calculation for softmax" (arXiv:1805.02867)
- Vaswani et al., "Attention Is All You Need" (NeurIPS 2017, arXiv:1706.03762)
- Dao-AILab/flash-attention 공식 저장소, 지원 GPU와 설치
'Inference > Optimization' 카테고리의 다른 글
| KV Cache 란 (0) | 2024.07.02 |
|---|---|
| prefill vs decode (0) | 2024.03.11 |
댓글