Skip to main content
ZICQ

Wiki Concepts

FlashAttention

Concepts
Aliases: Flash Attention FlashAttention-2 FlashAttention-3 ·2026-09-14

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.