All concepts
Flash Attention
Compute exact attention while reducing expensive GPU memory traffic.
Transformers & LLMs · Advanced · ~8 min
In plain English
Same maths as attention, better housekeeping: never write the giant score matrix to slow memory, compute it in tiles inside the fast memory instead.
Why it's worth your time
It made long contexts affordable without changing a single result — the rare optimization with no accuracy trade-off.
If you remember three things
- Exact, not approximate — identical output to standard attention
- It's an IO optimization, not a maths one
- Memory drops from quadratic to linear in sequence length
Overview
Exact attention computed without ever materializing the n×n score matrix in slow HBM. It tiles Q, K, and V into blocks that fit in on-chip SRAM and uses an online softmax to accumulate results, so memory scales O(n) instead of O(n²). This unlocks longer context and larger batches.
How it works
- Start: Q K V Blocks Attention normally materializes a huge n-by-n score matrix.
- Q K V Blocks -> Tiled SRAM FlashAttention processes blocks that fit in fast on-chip memory.
- Tiled SRAM -> Online Softmax It computes softmax statistics incrementally without storing the whole matrix.
- Online Softmax -> Exact Output The result is mathematically exact attention, but faster and more memory efficient.
- Exact Output -> Long Context Lower memory pressure makes longer sequences and larger batches practical.
In an interview
FlashAttention is an IO-aware exact attention algorithm. Instead of writing the full n×n scores to GPU HBM, it streams Q, K, and V in tiles through fast SRAM and keeps running softmax statistics (max and sum). It's bound by memory bandwidth, not capacity, giving the identical output with far less traffic.
Production defaults
- Turn it on
- always, if your hardware and library support it. There is no downside
- Expect
- 2–4× speedup on long sequences and a large memory reduction
- Check
- that your attention variant (ALiBi, sliding window) is supported by the kernel you're using
What breaks
- Results differ slightly from the reference — Floating-point accumulation order, not a bug. Differences should be at the noise floor — if they're larger, check the mask handling.
- No speedup — Short sequences are compute-bound, not memory-bound. The win scales with sequence length.