Flash attention and IO-aware attention
Be able to explain why flash attention is faster without changing the mathematics.
Prerequisites
- EContext length and the quadratic costrequired
- FThe memory hierarchy and bandwidthrequired
Intuition
Naive attention does this:
1. Compute S = QKᵀ → an N×N matrix, WRITE it to HBM
2. Read S, compute the softmax → another N×N matrix, WRITE it to HBM
3. Read P, compute PV → the output
For N = 8 192 every N×N matrix is 67 million elements — 134 MB in bf16. It is written and read several times, and that is where the time goes, not in the multiplications.
Flash attention does exactly the same mathematics but never moves the large matrix to HBM. Instead Q, K and V are split into blocks that fit in the SM's fast SRAM, and the whole computation is done there.
The result is identical — bit for bit, not just approximately. It is not an approximation. What changes is only in which order and in which memory the operations are carried out.
| Naive | Flash | |
|---|---|---|
| Memory | ||
| HBM traffic | where is the SRAM size | |
| Mathematics | exact | exact |
| Speed at N=8k | 1× | 2–4× |
Derivation
The problem to solve: a softmax normally requires seeing all the values before it can normalise. How do you compute it block by block?
The online softmax. Suppose you have already handled one block and have the running maximum , the running sum and the running output . A new block with values arrives:
- The new maximum:
- Correct the old sum, since it was computed with the wrong maximum:
- Correct the output in the same way:
The factor is the rescaling: everything computed with the old maximum is converted to the new one. Since the softmax is translation-invariant the result is exactly the same as if everything had been computed at once.
The backward pass has the same problem in reverse: saving (the large matrix) for the backward pass would reintroduce the memory. The solution is rematerialisation — save only and per row, and recompute block by block in the backward pass. It costs extra FLOPs but saves an enormous amount of memory traffic, and since the operation is memory-bound it is still faster.
What flash attention does NOT do:
| The misconception | The reality |
|---|---|
| «It is an approximation» | no, exactly the same result |
| «It changes the complexity to O(N)» | no, the FLOPs are still |
| «It makes long contexts free» | no, only cheaper in memory |
| «It makes every model faster» | only those where attention is the bottleneck |
The last is important: for short sequences the MLP blocks dominate, and flash attention then makes hardly any difference. The gain grows with the sequence length.
The versions: FlashAttention-2 improved the parallelisation over the sequence length and reached about 70 % of the theoretical maximum on an A100. FlashAttention-3 exploits the H100's asynchronous copying and fp8. In PyTorch it is built in as scaled_dot_product_attention, which chooses the backend automatically.
Code
import torch, torch.nn.functional as F, time
# In practice: use the built-in call, which chooses the flash backend automatically
def attention(q, k, v, causal=True):
return F.scaled_dot_product_attention(q, k, v, is_causal=causal)
# The online softmax — the core of the algorithm, in pure Python
def online_softmax_attention(Q, K, V, block_size=256):
"""The same result as ordinary attention, but without forming the whole N×N matrix."""
N, d = Q.shape
O = torch.zeros_like(Q)
for i in range(0, N, block_size):
q = Q[i:i + block_size]
m = torch.full((len(q), 1), float("-inf")) # the running maximum
l = torch.zeros((len(q), 1)) # the running sum
o = torch.zeros((len(q), d)) # the running output
for j in range(0, N, block_size):
kb, vb = K[j:j + block_size], V[j:j + block_size]
s = q @ kb.T / d ** 0.5
m_new = torch.maximum(m, s.max(dim=-1, keepdim=True).values)
corr = torch.exp(m - m_new) # rescale the old one
p = torch.exp(s - m_new)
l = corr * l + p.sum(dim=-1, keepdim=True)
o = corr * o + p @ vb
m = m_new
O[i:i + block_size] = o / l
return O
# Verify that the result is IDENTICAL, not approximate
torch.manual_seed(0)
N, d = 512, 64
Q, K, V = (torch.randn(N, d) for _ in range(3))
naive = torch.softmax(Q @ K.T / d ** 0.5, dim=-1) @ V
blocked = online_softmax_attention(Q, K, V, block_size=128)
print("the largest deviation:", float((naive - blocked).abs().max())) # ~1e-6, pure floating-point noise
# Memory use: N×N against O(N)
for N in (2048, 8192, 32768):
big_matrix_mb = N * N * 2 / 1024**2
print(f"N={N:>6}: the attention matrix {big_matrix_mb:>9.1f} MB in bf16")
# N= 2048: the attention matrix 8.0 MB in bf16
# N= 8192: the attention matrix 128.0 MB in bf16
# N= 32768: the attention matrix 2048.0 MB in bf16 ← per head, per layer
# The gain grows with the sequence length — measure it on your own hardware
def measure(N, d=64, heads=8, repeats=20):
dev = "cuda" if torch.cuda.is_available() else "cpu"
q, k, v = (torch.randn(1, heads, N, d, device=dev, dtype=torch.bfloat16)
for _ in range(3))
t0 = time.perf_counter()
for _ in range(repeats):
F.scaled_dot_product_attention(q, k, v, is_causal=True)
if dev == "cuda":
torch.cuda.synchronize()
return round((time.perf_counter() - t0) / repeats * 1000, 2)
The N=32768 row shows why this became necessary: two gigabytes per head and layer, just for an intermediate matrix that is thrown away immediately.
Mastery means
- Explains why naive attention is memory-bound
- Describes tiling and the online softmax
- Knows what flash attention does not change
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — arXiv (open access; licence per article)
- arXiv — FlashAttention-2 — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause