FlashAttention
FlashAttention is the attention algorithm rewrite proposed by Tri Dao et al. in 2022: keeping numerical accuracy, it reduces attention's HBM read/write to a minimum, achieving 3-5x speedup on long-sequence training/inference. Now standard in modern LLMs.
Core idea
Standard attention reads the Q/K/V matrices entirely from HBM (GPU memory) into SRAM/registers, writes intermediate results (row vectors for softmax) back to HBM, then writes back. This IO pattern is the bottleneck for long sequences.
FlashAttention's key insights:
- Tiling: split Q/K/V into tiles (blocks), compute each tile entirely in SRAM.
- Recomputation: on backward pass, recompute attention instead of storing intermediate results — trade compute for memory.
- Online softmax: chunked softmax, avoid storing the entire attention matrix (O(n²) memory).
Version history
- FlashAttention 1 (2022): 2-4x speedup, first version.
- FlashAttention 2 (2023): 5-9x speedup, better parallelism.
- FlashAttention 3 (2024): Hopper GPU (FP8) native support, ~2x additional speedup.
Performance numbers (H100, seq 8k)
| Implementation | Speed | Memory |
|---|---|---|
| Standard PyTorch | 1x baseline | 1x baseline |
| FlashAttention 2 | ~5x | ~10x savings |
| FlashAttention 3 | ~10x | ~20x savings |
Practical deployment
- PyTorch 2.x: built-in
torch.nn.functional.scaled_dot_product_attention, auto-uses FlashAttention. - HuggingFace Transformers:
attn_implementation="flash_attention_2". - vLLM / TGI: enabled by default in inference.
Limitations
- Requires CUDA (H100/A100/4090 etc), AMD / Apple Silicon not supported.
- Some older GPUs lack compute (needs SM 8.0+, A100+).
- Numerically doesn't change attention output, just faster and cheaper.