Source-linked AI summary
FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré
TL;DR
Long-context Transformer는 self-attention의 quadratic time 및 memory cost로 제약된다. FlashAttention은 IO-aware tiling으로 GPU memory traffic을 줄여 HBM access를 최대 9× 감소시키면서 더 높은 품질의 longer-context model을 지원한다.
문제
Long-context Transformer는 quadratic self-attention의 time 및 memory cost에 직면해 더 빠르고 memory-efficient한 attention이 요구된다.
방법
FlashAttention은 IO-aware tiling과 recomputation으로 exact attention을 계산하고 GPU HBM과 on-chip SRAM 사이의 read 및 write를 줄인다.
결과
Standard attention보다 HBM access가 최대 9× 적으며, 모든 SRAM size에서 asymptotically FlashAttention을 개선하는 exact algorithm은 없다.
시사점 및 한계
FlashAttention은 더 긴 context를 지원하는 Transformer를 가능하게 하며, Path-X와 Path-256에서 chance보다 나은 성능을 달성하고 model quality를 향상한다.
시사점 및 한계
IO-aware attention 구현에는 상당한 CUDA engineering effort가 필요하며 GPU architecture 간에 그대로 이전되지 않을 수 있다.
Abstract
from arXiv · showhide
Transformers are slow and memory-hungry on long sequences, since the time and memory complexity of self-attention are quadratic in sequence length. Approximate attention methods have attempted to address this problem by trading off model quality to reduce the compute complexity, but often do not achieve wall-clock speedup. We argue that a missing principle is making attention algorithms IO-aware -- accounting for reads and writes between levels of GPU memory. We propose FlashAttention, an IO-aware exact attention algorithm that uses tiling to reduce the number of memory reads/writes between GPU high bandwidth memory (HBM) and GPU on-chip SRAM. We analyze the IO complexity of FlashAttention, showing that it requires fewer HBM accesses than standard attention, and is optimal for a range of SRAM sizes. We also extend FlashAttention to block-sparse attention, yielding an approximate attention algorithm that is faster than any existing approximate attention method. FlashAttention trains Transformers faster than existing baselines: 15% end-to-end wall-clock speedup on BERT-large (seq. length 512) compared to the MLPerf 1.1 training speed record, 3$\times$ speedup on GPT-2 (seq. length 1K), and 2.4$\times$ speedup on long-range arena (seq. length 1K-4K). FlashAttention and block-sparse FlashAttention enable longer context in Transformers, yielding higher quality models (0.7 better perplexity on GPT-2 and 6.4 points of lift on long-document classification) and entirely new capabilities: the first Transformers to achieve better-than-chance performance on the Path-X challenge (seq. length 16K, 61.4% accuracy) and Path-256 (seq. length 64K, 63.1% accuracy).
1 서론
FlashAttention은 tiling을 통해 GPU HBM 접근을 줄이는 IO-aware exact attention 알고리즘을 도입하며, 더 빠른 long-context modeling을 위해 block-sparse attention으로 확장된다. 논문은 더 빠른 학습, 향상된 품질, 최대 64K 길이의 sequence로의 확장을 보고한다.
- FlashAttention: FlashAttention은 tiling과 online softmax reduction을 사용해 HBM에 큰 attention matrix를 materialize하지 않으면서 exact attention을 계산한다.또한 backward pass를 위해 큰 intermediate attention matrix를 저장하지 않아 memory access를 줄인다.
- IO Complexity: FlashAttention은 standard attention의 Ω(Nd + N^2)에 비해 O(N^2d^2M^-1)의 HBM access를 필요로 하며, 최대 9× fewer access와 exact attention의 asymptotic lower bound를 달성한다.여기서 d는 head dimension이고 M은 SRAM size다. 이 알고리즘은 SRAM size의 일정 범위에서 optimal하다.
- Block-Sparse FlashAttention: Block-sparse FlashAttention은 IO-aware primitive를 approximate attention에 적용해 FlashAttention보다 2–4× speedups를 달성하면서 sequence length 64K까지 확장된다.이 접근법은 여러 sparse 및 low-rank approximate attention method를 제한하는 memory-access overhead를 해결한다.
- Faster Model Training: FlashAttention은 baseline implementation보다 BERT-large를 15% faster, GPT-2를 3× faster, long-range arena model을 2.4× faster하게 학습한다.보고된 설정은 각각 sequence length 512, 1K, 1K–4K를 사용한다.
- Higher Quality Models: Longer-context modeling은 GPT-2 perplexity를 0.7 개선하고, long-document classification을 6.4 points 향상시키며, sequence length 16K에서 better-than-chance Path-X performance를 가능하게 한다.Block-sparse FlashAttention은 이러한 capability를 sequence length 64K까지 확장하며, 이때 Path-256은 63.1% accuracy에 도달한다.
- Benchmarking Attention: FlashAttention은 sequence length 128–2K에서 standard attention보다 최대 3× faster하고, 64K까지 확장되며, block-sparse FlashAttention은 기존 approximate method보다 빠르다.Sequence length 512까지 FlashAttention은 기존 attention method보다 빠르고 memory-efficient하다. 1K를 넘으면 일부 approximate method가 더 빨라진다.
2 배경
GPU 성능은 더 작은 메모리일수록 빠르지만 용량은 훨씬 작아지는 메모리 계층 전반의 연산과 메모리 접근 간 균형에 좌우된다. Standard attention은 HBM에 O(N^2) 크기의 중간 결과를 구체화하고 이를 반복해서 접근하므로 느리고 메모리를 많이 사용한다.
- GPU 메모리 계층: GPU 메모리는 계층적이다. A100 HBM은 40–80GB 용량과 1.5–2.0TB/s 대역폭을 제공하는 반면, 온칩 SRAM은 multiprocessor당 192KB이며 약 19TB/s의 대역폭을 제공한다.SRAM은 HBM보다 대략 한 자릿수 배 빠르지만 용량은 수십 자릿수 배 작다.
- 성능 특성: Arithmetic intensity는 compute-bound 연산과 memory-bound 연산을 구분한다. 후자의 실행 시간은 연산보다 메모리 접근에 의해 지배된다.softmax, dropout, normalization 같은 elementwise 연산과 reduction은 대표적인 memory-bound 사례다.
- Kernel fusion: Kernel fusion은 공유 입력을 한 번만 로드해 HBM 트래픽을 줄이지만, 학습에서는 backward pass를 위해 중간 값을 여전히 HBM에 기록해야 한다.이 요구사항 때문에 compiler가 많은 elementwise 연산의 fusion을 지원하더라도 naive fusion의 효과는 제한된다 [53] [65].
- Standard attention: Standard attention은 HBM에 S와 P를 구체화하므로 O(N^2) memory와 많은 메모리 접근이 필요하며, 특히 N ≫ d일 때 wall-clock time이 느려진다.구현은 S를 기록하고 softmax를 위해 다시 읽은 다음, P를 기록하고 V와 함께 다시 읽어 O를 생성한다.
3 FlashAttention: 알고리즘, 분석 및 확장
FlashAttention은 tiling, online softmax normalization, recomputation, fused GPU kernels를 결합해 sub-quadratic HBM access로 exact attention을 계산한다. IO analysis는 SRAM size 전반에서 점근적 최적성을 확립하며, block-sparse FlashAttention은 sparsity에 비례해 IO를 줄인다.
- 알고리즘: Tiling은 Q, K, V block을 SRAM에 로드하고 normalization statistics를 사용해 blockwise softmax 결과를 결합하며, kernel fusion을 통해 반복적인 HBM transfer를 피한다.Recomputation은 저장하는 대신 backward pass에서 attention matrix를 재구성한다.
- 알고리즘: FlashAttention은 O(N^2d) FLOPs로 exact attention을 수행하며 input과 output 외에 O(N) additional memory만 사용한다.Output과 softmax statistics를 유지해 backward-pass recomputation에 활용하므로 O(N^2) intermediate matrix를 저장하지 않는다.
- IO Complexity: FlashAttention은 standard attention의 Θ(Nd + N^2)에 비해 Θ(N^2d^2M^-1) HBM access를 필요로 하며, d=64–128이고 M이 약 100KB일 때 often many fewer accesses를 달성한다.HBM access 감소는 backward recomputation으로 FLOP count가 증가하더라도 더 빠른 실행과 더 작은 memory footprint를 설명한다.
- IO Complexity: 어떤 exact attention algorithm도 모든 SRAM size M ∈ [d, Nd]에 대해 o(N^2d^2M^-1) HBM access를 달성할 수 없다.따라서 FlashAttention은 제시된 SRAM range 전반에서 IO-optimal이지만, M에 parameterized된 lower bound는 여전히 미해결 상태다.
- Block-Sparse Extension: Block-sparse FlashAttention은 zero attention block을 건너뛰며, nonzero block의 fraction을 s라 할 때 Θ(Nd + N^2d^2M^-1s) HBM access를 필요로 한다.IO complexity에서 더 큰 항을 sparsity factor s만큼 개선한다.
4 실험
FlashAttention은 모델 품질을 바꾸지 않으면서 Transformer 학습 속도를 높이고 더 긴 context를 지원하며, attention 실행 시간과 메모리 사용량을 줄인다. 이러한 개선으로 long-context 품질이 높아지고 Path-X와 Path-256에서 chance 수준을 넘는 성능을 달성한다.
- 학습 속도: FlashAttention을 사용하면 BERT-large 학습은 15% 빠르고, GPT-2 학습은 최대 3× 빠르며, Long-Range Arena 학습은 2.4× 빠르다.GPT-2 비교에서 HuggingFace 대비 최대 3×, Megatron-LM 대비 최대 1.7× 빠르며, BERT는 동일한 초기화에서 같은 목표 정확도에 도달한다.
- 실험적 검증: FlashAttention은 model definition을 바꾸지 않으므로 baseline perplexity와 학습 곡선을 유지하지만, 재현된 Long-Range Arena baseline은 tuning에 크게 의존한다.Long-Range Arena 실험은 1,024부터 4,096까지의 sequence length를 사용하며 accuracy, throughput, training time을 보고한다.
- Long Context를 활용한 언어 모델링: 1K context를 사용하는 Megatron-LM 대비 4K context에서 GPT-2 학습이 30% 빠르며, perplexity도 0.7 더 좋다.FlashAttention은 최적화된 Megatron-LM 구현보다 빠른 상태를 유지하면서 GPT-2 context length를 4× 늘린다.
- 장문서 분류: sequence를 늘리면 length 512 대비 MIMIC-III 16K에서 document classification이 4.3 points, ECtHR 8K에서 8.5 points 향상된다.두 dataset은 매우 긴 문서로 구성되며, 최대 길이는 MIMIC-III가 14,562 tokens, ECtHR이 49,392다.
- Path-X와 Path-256: FlashAttention은 Path-X에서 61.4% accuracy를 달성하고, block-sparse FlashAttention은 sequence length 64K의 Path-256에서 63.1%에 도달한다.이는 두 난해한 long-context task에서 Transformer가 non-random performance를 달성한 최초의 결과로 보고된다.
- Attention 벤치마킹: FlashAttention은 PyTorch exact attention보다 최대 3× 빠르고 최대 20× 메모리 효율적이며, 메모리는 sequence length에 따라 선형적으로 증가한다.Block-sparse FlashAttention은 동일한 memory footprint를 가지며, 실행 시간도 선형적으로 확장되고 기존 approximate-attention baseline을 능가한다.
5 한계와 향후 방향 … B.3 FlashAttention: Forward Pass
이 논문은 memory-efficient forward 및 backward attention 연산과 FlashAttention의 tiled GPU 구현을 개발하고, IO-aware deep learning 및 multi-GPU 시스템의 한계와 확장 방향을 제시한다. 또한 efficient Transformer, structured-matrix, sparse-training, runtime-optimization 연구 흐름 속에서 이 접근법의 위치를 설명한다.
- 5 한계와 향후 방향: 이 접근법의 주요 한계는 engineering complexity다. 새로운 IO-aware attention 구현마다 GPU 아키텍처 간에 이식되지 않을 수 있는 새로운 low-level CUDA kernel이 필요하다.저자들은 attention 알고리즘을 표현하기 위한 향후 방향으로 PyTorch와 같은 high-level language를 제안한다.
- 5 한계와 향후 방향: 저자들은 모든 deep-learning layer가 GPU HBM에 접근하므로 IO-aware optimization을 attention을 넘어 확장할 것을 제안하며, multi-GPU data transfer를 추가적인 IO-analysis 문제로 식별한다.단일 GPU에서는 attention 구현이 상수항을 제외하면 최적이지만, multi-GPU parallelization에는 GPU 간 transfer를 고려해야 한다고 설명한다.
- A 관련 연구: 이 연구는 IO-aware optimization을 I/O complexity, working-set models, data locality, Roofline arithmetic intensity, scalability analyses 와 연결한다.이러한 연결은 memory-hierarchy optimization이 attention에 특화된 아이디어가 아니라 오래된 systems principle임을 보여준다.
- A 관련 연구: 이 논문은 efficient attention과 block-sparse variant를 structured matrices, sparse training, 그리고 quadratic sequence-length costs를 다루는 대안적 long-context Transformer module과 연관 짓는다.예로 sparse 및 low-rank matrices, pruning과 lottery-ticket training, 그리고 Reformer [51], Performer [12] [54], S4 [31] [36] [37], FLASH [42]와 같은 모델을 들 수 있다.
- B 알고리즘 세부 사항: FlashAttention은 linear extra memory를 사용해 memory-efficient forward 및 backward attention을 도출하고, runtime과 memory footprint를 개선하도록 HBM accesses를 줄인다.Naive memory-efficient formulation은 quadratic intermediate storage를 피하지만 여전히 quadratic HBM accesses를 발생시킨다. FlashAttention은 이러한 accesses를 직접 해결한다.
- B.1 Memory-efficient forward pass: Forward pass는 quadratic attention matrix를 저장하지 않고 softmax normalization constants를 별도로 유지하면서 output contributions를 반복적으로 누적한다.이를 통해 O(n) extra memory를 사용하며, normalization constants에 O(n) memory를, 각 output computation에 O(d) extra memory를 사용한다.
- B.2 Memory-efficient backward pass: Backward pass는 gradient checkpointing에만 의존하지 않고 linear memory로 dV, dQ, dK를 명시적으로 계산한다.유도는 dP와 dS를 거쳐 진행되며, L과 D 같은 linear-size intermediate만 저장한다.
- B.3 FlashAttention: Forward Pass: FlashAttention의 forward pass는 HBM과 on-chip SRAM 사이에서 Q, K, V를 tiles로 나누고, chip에서 masked 및 dropout-adjusted attention을 계산하며, backpropagation을 위해 softmax statistics를 저장한다.알고리즘은 Q를 row blocks로, K,V를 column blocks로 partition한 뒤 O, ℓ, m, 그리고 random-number-generator state R을 반환한다.
B.4 FlashAttention: backward pass · B.5 Rabe and Staats [66]와의 비교
FlashAttention의 backward pass는 O(N^2) FLOPs와 O(N) 추가 메모리로 exact gradient를 유지하면서 standard attention보다 HBM access를 줄인다. Rabe and Staats [66]와 비교하면 memory-access reduction을 목표로 하며, 상당한 memory savings를 유지하면서 더 높은 속도를 달성한다.
- B.4 FlashAttention: backward pass: FlashAttention은 SRAM에서 backward gradient를 blockwise로 계산하며, full attention matrix를 materialize하지 않고 attention과 dropout 정보를 재구성하면서 dQ, dK, dV를 생성한다.softmax-gradient computation은 SRAM에 들어가지 않을 수 있는 N-sized row에 대한 reduction을 피하고, 동등한 reformulation을 사용한다.
- B.4 FlashAttention: backward pass: O(N^2) FLOPs와 O(N) extra memory는 input, output, output gradient, input gradient를 제외한 FlashAttention backward pass의 특성이다.backward pass는 pseudo-random generator state를 저장하고 replay하여 O(N^2) dropout mask 저장을 피한다.
- B.4 FlashAttention: backward pass: d ≤ M ≤ Nd일 때 FlashAttention backward의 HBM access는 Θ(N^2d^2M^-1)이고, standard attention은 Θ(Nd + N^2)이다.이는 sequence length N, head dimension d, SRAM size M에 대한 IO-complexity 비교다.
- B.5 Rabe and Staats [66]와의 비교: FlashAttention과 Rabe and Staats [66]는 모두 attention block을 tile하고, 큰 forward attention matrix 저장을 피하며, backward pass에서 이를 재계산한다.두 방법 모두 tiling 또는 softmax scaling [51] [60]을 사용한다.
- B.5 Rabe and Staats [66]와의 비교: FlashAttention은 memory access를 줄이는 반면, Rabe and Staats [66]는 주로 maximum GPU-memory footprint를 줄인다. memory access가 주요 runtime determinant로 식별된다.memory access를 줄이면 total memory requirement도 필연적으로 감소한다.
- B.5 Rabe and Staats [66]와의 비교: standard attention보다 2-4× 빠른 FlashAttention과 달리, Rabe and Staats [66]는 standard-attention speed 수준이거나 약간 더 느리지만, 두 방법 모두 상당한 memory를 절약한다.이 비교는 FlashAttention의 속도 우위가 total memory footprint만이 아니라 memory access를 줄이는 데서 비롯된다고 설명한다.
- B.5 Rabe and Staats [66]와의 비교: FlashAttention은 block output을 incrementally update하므로, temporary output을 저장한 뒤 normalization statistic으로 결합하는 Rabe and Staats [66]보다 total memory가 적게 필요하다.이 차이는 attention block 전반에서 정보를 요약하고 전파하는 방식에 관한 것이다.
- B.5 Rabe and Staats [66]와의 비교: FlashAttention은 backward computation을 analytical하게 단순화하여 각 block의 temporary output이 아니라 attention matrix만 recompute하며, 이로써 memory requirement를 줄이고 Rabe and Staats [66]보다 speedup을 얻는다.반면 Rabe and Staats [66]는 gradient checkpointing을 사용해 두 quantity를 모두 recompute한다.
C 증명 · D 확장 세부사항 · D.1 Block-sparse FlashAttention
증명은 FlashAttention이 exact attention을 계산하면서 forward 및 backward pass의 HBM access를 줄인다는 점을 보이며, block-sparse extension에서는 nonzero block의 비율에 따라 access 규모가 조정됨을 확립한다.
- C 증명: forward algorithm은 exact하다. 최종 outer-loop iteration 이후 출력은 softmax(QK^T)V와 같다.정확성은 row maximum, exponential row sum, 누적 output을 유지하면서 key-value block에 대해 induction을 적용해 따른다.
- C 증명: FlashAttention의 forward pass는 Θ(NdT_c) HBM access를 요구한다. 각 K 및 V element는 한 번 load되는 반면 Q와 O는 T_c개 pass에 걸쳐 다시 참조되기 때문이다.반면 standard attention은 Θ(Nd + N^2) global-memory access를 요구한다.
- C 증명: IO bound는 on-chip SRAM에 들어맞도록 K/V, Q/O, score block을 선택하는 데 좌우된다.증명은 B_c×d, B_r×d, B_r×B_c 크기의 block에 대한 제약을 각각 도출한다.
- C 증명: exact-attention IO lower bound는 Q, K, V, O가 합쳐서 최소 Ω(Nd) HBM access를 요구하기 때문에 성립한다.이는 M = Θ(Nd) regime에서 access count가 해당 input-output cost보다 작아지는 어떠한 exact algorithm도 반박한다.
- C 증명: FlashAttention의 backward pass 역시 Θ(NdT_c) HBM access를 요구하며, K와 V를 한 번 load하고 dK와 dV를 한 번 write한다.Standard attention의 backward는 Θ(Nd + N^2) HBM access를 요구한다.
- D 확장 세부사항: block-sparse extension은 FlashAttention의 tiled computation을 유지하면서 nonzero sparsity mask가 선택한 block만 load한다.Algorithm 5는 각 mask block을 검사하고 M_ij ≠ 0인 경우에만 on-chip computation을 수행한다.
- D.1 Block-sparse FlashAttention: Block-sparse FlashAttention은 zero block을 건너뛰어 nonzero-block fraction s에 비례해 HBM access를 줄이면서도 N×d output은 계속 write한다.그 외에는 algorithm이 FlashAttention과 동일하며, s가 작을 때 IO reduction은 O를 write하는 cost에 의해 제한된다.
D.2 잠재적 확장
IO-aware 접근법은 single-GPU attention을 넘어 multi-GPU 협력, sparse MLP 최적화, kernel machine learning을 포함한 확장을 제시한다. 이러한 방향은 메모리 계층 비대칭성 또는 저차원 입력에서 반복되는 계산을 활용해 메모리 및 계산 비용을 줄인다.
- Multi-GPU Attention: Multi-GPU attention은 여러 메모리 계층을 고려하면서 매우 긴 시퀀스에 대해 동일한 노드의 GPU들이 협력하도록 할 수 있다 [77].계층은 GPU SRAM, 로컬 GPU HBM, 다른 GPU의 HBM으로 구성되며, attention은 일반적으로 4–8개의 GPU에 분할된다.
- Sparse MLP 계층: IO-aware 구현은 메모리 트래픽으로 인해 sparsity에 비례하는 속도 향상이 어려운 경우 sparse MLP 계층을 더 효율적으로 만들 수 있다 [17].제안된 방향은 대규모 모델의 계산 요구량을 줄이는 것을 목표로 한다.
- Kernel machine learning: Kernel machine learning에서도 유사한 기회를 활용할 수 있다. 각 N×N kernel entry가 차원 d ≪ N인 두 벡터에 의존하므로 입력을 반복해서 로드하고 재계산할 수 있기 때문이다.FlashAttention은 QK^⊤의 유사한 low-rank 구조를 활용해 필요한 attention block을 재계산함으로써 HBM access를 줄인다.
E 전체 실험 결과 · E.1 BERT
BERT-large 실험은 MLPerf 1.1 절차를 따르며, 동일한 평가 조건에서 Nvidia가 보고한 제출 결과와 training speed를 비교한다. Training은 8×A100-80GB GPUs에서 FP16으로 수행하고, runtime은 10회 실행의 평균으로 산출한다.
- E.1 BERT: BERT-large training은 LAMB optimization, batch size 448, 최대 7100 steps를 포함한 reference MLPerf 1.1 procedure와 hyperparameters를 따른다.learning rate는 3.75e-3이다.
- E.1 BERT: masked-language-modeling validation accuracy가 72.0%에 도달하면 training을 중단하고, 이후 wall-clock runtime을 측정한다.실행은 Apex AMP를 O2 optimization level로 사용해 FP16 precision에서 수행한다.
- E.1 BERT: 실험은 training speed를 MLPerf 1.1에 제출된 Nvidia의 보고 결과와 비교한다.비교 결과는 Table 1에 제시한다.
- E.1 BERT: 평가는 MLPerf 1.1 reference implementation과 동일한 train/validation split을 사용하며, Nvidia baseline과 동일한 10,000 validation examples를 포함한다.
- E.1 BERT: model은 8×A100-80GB GPUs에서 training한다.
- E.1 BERT: 각 training run에는 16–19분이 걸리며, 보고 결과는 10 runs의 평균이다.
E.2 GPT-2 · E.3 LRA 세부 사항
GPT-2 실험은 동일한 학습 조건에서 표준 Hugging Face 및 Megatron-LM 구현을 사용했으며, LRA 비교는 공개된 설정을 따랐고 태스크 전반에서 튜닝 후 정확도가 유사했다. FlashAttention은 baseline GPT-2 validation perplexity 곡선과 일치했으며, 평가 절차에는 재현성과 구현상의 제약을 반영했다.
- E.2 GPT-2: GPT-2는 Megatron-LM training recipe 를 따라 표준 Hugging Face 및 Megatron-LM 구현을 사용했다.
- E.2 GPT-2: GPT-2 모델은 512 effective batch size, AdamW, 모델별 learning rate, weight decay 0.1, 400K steps 및 mixed-precision training을 사용했다.learning rate는 GPT-2 small에서 6e-4, GPT-2 medium에서 1.5e-4였다.
- E.2 GPT-2: GPT-2 training에는 GPT-2 BPE tokenizer를 적용한 OpenWebText와 모델 간 공유되는 고정된 무작위 0.5% validation split을 사용했다.
- E.2 GPT-2: GPT-2 wall-clock training은 8×A100-40GB GPU에서 측정했으며, small 모델은 2.7–9.5 days, medium 모델은 6.9–21.0 days가 걸렸다.보고된 기간은 Table 2에 해당한다.
- E.2 GPT-2: FlashAttention은 GPT-2 small 및 medium에서 Hugging Face baseline과 거의 동일한 validation perplexity 곡선을 생성했다.비교 결과는 Figure 4에 제시되어 있으며, 구현 간 validation 동작이 동등함을 보여준다.
- E.3 LRA 세부 사항: LRA 실험은 Long-range arena paper와 repository 및 Nyströmformer reproduction [80]의 설정을 따랐다.reproduction 성능이 낮을 경우 Tay et al. [80] 또는 Xiong et al. 이 보고한 더 나은 baseline 결과를 사용했다.
- E.3 LRA 세부 사항: 거의 모든 attention method는 hyperparameter tuning 후 5개 LRA task 전반에서 유사한 정확도를 달성했다.
- E.3 LRA 세부 사항: LRA overall wall-clock speedup은 5개 task별 speedup의 geometric mean으로 계산했다.불안정했던 Performer와 구현에 FP16 support가 없었던 Local Attention을 제외하고 mixed precision을 사용했다.
E.4 Apex FMHA와의 비교
FlashAttention은 tiling과 recomputation을 적용해 memory 사용량을 줄이고 더 긴 sequence와 더 폭넓은 hardware를 지원하면서도 짧은 sequence에서 comparable한 runtime을 달성한다. attention matrix를 저장하는 대신 recompute하기 때문에 forward pass에서는 약간 빠르지만 backward pass에서는 약간 느리다.
- Apex FMHA와의 비교: Apex FMHA는 head dimension이 64인 BERT model을 대상으로 하며 A100 GPU와 sequence length 최대 512만 지원한다.dropout(softmax(mask(QK^T)))V를 하나의 CUDA kernel로 fuse하지만 gradient 계산을 위해 attention matrix를 HBM에 저장하므로 memory 절감이 제한된다.
- Apex FMHA와의 비교: FlashAttention은 최대 64K의 sequence, 16, 32, 64, 128의 head dimension, 그리고 작성 시점에 설명된 모든 Turing 및 Ampere GPU를 지원한다.이러한 확장은 tiling과 recomputation을 사용해 긴 sequence를 처리하고 memory를 절약한다.
- Apex FMHA와의 비교: FlashAttention은 일반적으로 forward pass에서 FMHA보다 약간 빠르고 backward pass에서는 약간 느리다.backward pass의 slowdown은 forward pass 동안 attention matrix를 저장하지 않고 backpropagation 중에 recompute하기 때문에 발생한다.
- Apex FMHA와의 비교: FlashAttention은 sequence length 128에서 FMHA보다 약 4% 느리다.비교는 batch size 64, 16개 head, head dimension 64인 A100-SXM4-40GB GPU에서 masking과 dropout을 적용해 수행한다.
E.5 다양한 하드웨어 및 구성에서의 속도 향상
FlashAttention의 속도 향상은 GPU 하드웨어와 구성에 따라 달라진다. HBM 대역폭과 SRAM 크기가 IO 효율에 영향을 주기 때문이다. 테스트한 설정 전반에서 상당한 성능 향상을 보이며, masking과 dropout은 kernel fusion을 통해, 낮은 메모리 대역폭은 대체로 속도 향상을 키우는 반면, 큰 head dimension과 작은 SRAM은 이를 줄인다.
- A100: A100에서는 모든 sequence length에서 표준 PyTorch attention 대비 2-4× speedup이 나타나며, kernel fusion을 통한 dropout과 masking 적용으로 증가한다.이 구성은 batch size 8, head dimension 64, attention heads 12개를 사용한다.
- RTX 3090: RTX 3090에서 2.5-4.5× speedup이 발생하며, 메모리 대역폭이 더 낮기 때문에 A100을 약간 웃돈다. 메모리 대역폭은 대략 900 GB/s 대 1.5 TB/s다.RTX 3090 구성은 batch size 12와 12개의 attention heads를 사용한다.
- T4: T4는 SRAM이 더 작아 FlashAttention 블록을 더 작게 설정해야 하므로 A100보다 less speedup을 보이며, 이는 IO complexity analysis와 일치한다.T4 GPU는 일반적으로 inference에 사용되므로, combined forward-plus-backward와 forward-only speedup을 모두 보고한다.
E.6 전체 벤치마킹 결과
이 절에서는 다양한 sequence length에서 A100을 사용해 exact, approximate, sparse attention 구현을 벤치마크하고, 여러 조건에서 forward, backward, combined runtime과 memory usage를 보고한다. 비교 대상에는 Reformer [51], Local Attention, Linformer [84], Smyrf, LongShortFormer (LSFormer), Block-Sparse Attention [11], Longformer [3], BigBird 등 FlashAttention 관련 baseline이 포함된다.
- E.6 전체 벤치마킹 결과: 벤치마크에서는 PyTorch/HuggingFace, Megatron, Reformer [51], Local Attention, Linformer Attention [84], Smyrf, LongShortFormer (LSFormer), Block-Sparse Attention [11], Longformer [3], BigBird 의 reference implementation과 exact, approximate, sparse attention을 비교한다.이 baseline들은 주요 exact, approximate, sparse attention 범주를 포괄한다.
- E.6 전체 벤치마킹 결과: 실험에서는 dimension이 64인 head 8개, batch size 16, random Q, K, V vector를 사용하고, 하나의 40 GB A100 GPU에서 sequence length를 변화시켰다.Hidden layer에서의 attention projection은 제외했으며, dropout은 0.1이고 masking에는 uniformly random mask length를 사용했다.
- E.6 전체 벤치마킹 결과: 결과 표에는 unconditioned, masked, dropout, combined dropout-and-masking 설정에서 sequence length에 따른 runtime과 memory usage가 제시된다.표에서는 각각 bold와 밑줄 서식을 사용해 최선 및 차선 method를 식별한다.
- E.6 전체 벤치마킹 결과: Runtime은 dropout, masking 또는 둘 다 적용하거나 적용하지 않은 forward, backward, combined pass에 대해 보고하며, memory는 어느 것도 적용하지 않은 combined forward 및 backward pass에 대해 측정한다.Local Attention은 FP32만 지원하는 구현이므로 이를 제외하면 측정에는 FP16을 사용했다.
- E.6 전체 벤치마킹 결과: Sequence length는 GPU memory가 소진될 때까지 증가시키되, implementation limit을 적용했다. Megatron은 2048, Block-Sparse는 4096, Longformer와 BigBird는 8092다.외부 library의 bug로 해당 실행이 불가능했기 때문에 Block-Sparse, Longformer, BigBird는 masked backward pass 없이 측정했다.