Fast Attention and Long Context
FlashAttention
FlashAttention gives the same answer as ordinary attention while never writing the giant score table to memory, which is what makes long contexts affordable.
- 14 min read
- 3 reading levels
- Updated
Read these first
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
FlashAttention is the attention you already know, rearranged so the computer never writes down the enormous middle step.
The kitchen you have cooked in
Picture a small kitchen. The fridge is across the room and the counter is right in front of you. Fetch one onion, walk back, chop it, then walk to the fridge again. You will spend your evening walking.
A good cook carries a whole tray over once. Everything happens on the counter, and only the finished dish goes back.
A graphics card has the same two places. It has a large slow store, and a tiny fast workspace beside each processing unit. FlashAttention is the cook who stopped walking.
What was broken before
Attention compares every word with every other word. Attention here means the step where the model decides which earlier words matter for the word it is writing now. If you have not met it yet, read attention first.
The ordinary way writes that comparison out in full. For a thousand words that is a table with a million entries. For a hundred thousand words it is ten billion entries.
That table is written to the slow store and read back. Nobody wants the table. It exists for a moment and is thrown away.
So the model spends most of its time moving a throwaway table around. The chopping is fast. The walking is slow.
The trick, in one idea
You do not need the whole table to get the answer. You need a running total — a summary you keep updating as you go. It is like adding up a shopping bill in your head.
FlashAttention brings over one strip of the table at a time. It updates the running total on the fast counter. Then it throws the strip away and fetches the next one.
the old way FlashAttention
----------- --------------
compute the WHOLE table take one strip
write it to the slow store work on it on the fast counter
read it back update a running total
finish throw the strip away, take the next
finish
answer: identical answer: identical
memory: enormous memory: tinyThere is one wrinkle that makes this harder than a shopping bill. Attention turns scores into percentages, and a percentage depends on every score, including ones you have not seen yet.
This part is confusing for almost everyone the first time. Read it twice, that is normal.
The fix is a correction step. A later strip may contain a bigger score than anything seen so far. The running total is then scaled down to match. The bookkeeping is exact, so the final answer is not an approximation.
Why the word "exact" matters
Many long-context methods drop comparisons to save time. They give an answer that is close, not equal.
FlashAttention drops nothing. The output is the same output, to the last decimal the hardware can hold. That is why it spread everywhere within a year, with no quality argument to have.
Where you have already seen it
- Every chatbot that accepts a long document you paste in.
- Fine-tuning a model on a free cloud notebook without running out of memory.
- Any local model server running on your own machine.
Remember this
- FlashAttention never stores the full comparison table. It works one strip at a time.
- The saving is in memory traffic, meaning data carried between the slow store and the fast workspace.
- The answer is exact, not approximate. Nothing is skipped.
What to learn next
- Memory-bound vs compute-bound — why moving data, not multiplying, sets the clock.
- Attention — the operation FlashAttention makes cheap, from first principles.
- Context window — what a longer context buys you, and what it costs.
Developer — Code and libraries.
Setup
pip install numpy # for the first example
pip install torch # for the second; a CUDA GPU is needed for the measurementWritten and run against NumPy 1.26.4 and PyTorch 2.5.1+cu121 on Python 3.10.
The algorithm, in plain NumPy
The whole idea is the loop below. It keeps three running quantities and never builds an N x N array.
import numpy as np
rng = np.random.default_rng(0)
N, d, B = 512, 32, 128 # 512 tokens, head dim 32, tiles of 128 keys
Q = rng.normal(size=(N, d)).astype(np.float32)
K = rng.normal(size=(N, d)).astype(np.float32)
V = rng.normal(size=(N, d)).astype(np.float32)
scale = 1.0 / np.sqrt(d)
def naive_attention(Q, K, V):
S = (Q @ K.T) * scale # the N x N score matrix, held in full
S = S - S.max(axis=-1, keepdims=True)
P = np.exp(S)
return (P / P.sum(axis=-1, keepdims=True)) @ V, S.nbytes
def flash_attention(Q, K, V, block=B):
n = Q.shape[0]
O = np.zeros_like(Q) # running output
m = np.full(n, -np.inf, dtype=np.float32) # running row max
l = np.zeros(n, dtype=np.float32) # running sum of exp
biggest = 0
for j in range(0, n, block): # one tile of keys at a time
Kj, Vj = K[j:j + block], V[j:j + block]
Sj = (Q @ Kj.T) * scale # N x block, never N x N
biggest = max(biggest, Sj.nbytes)
m_new = np.maximum(m, Sj.max(axis=-1))
correction = np.exp(m - m_new) # rescale what we already accumulated
Pj = np.exp(Sj - m_new[:, None])
l = correction * l + Pj.sum(axis=-1)
O = correction[:, None] * O + Pj @ Vj
m = m_new
return O / l[:, None], biggest
out_naive, bytes_naive = naive_attention(Q, K, V)
out_flash, bytes_flash = flash_attention(Q, K, V)
print("outputs match:", np.allclose(out_naive, out_flash, atol=1e-5))
print("largest absolute difference:", float(np.abs(out_naive - out_flash).max()))
print()
print(f"naive: biggest score matrix held = {bytes_naive:,} bytes ({N} x {N})")
print(f"flash: biggest score tile held = {bytes_flash:,} bytes ({N} x {B})")
print(f"ratio = {bytes_naive / bytes_flash:.0f}x smaller")
print()
for n in (1024, 8192, 131072):
full = n * n * 2 / 1e9 # bf16 scores for one head
print(f"seq {n:>7}: one head's full score matrix in bf16 = {full:10.2f} GB")outputs match: True largest absolute difference: 2.384185791015625e-07 naive: biggest score matrix held = 1,048,576 bytes (512 x 512) flash: biggest score tile held = 262,144 bytes (512 x 128) ratio = 4x smaller seq 1024: one head's full score matrix in bf16 = 0.00 GB seq 8192: one head's full score matrix in bf16 = 0.13 GB seq 131072: one head's full score matrix in bf16 = 34.36 GB
Reading that output
The difference is 2.4e-07. That is float32 rounding, not an approximation. This is the claim that made FlashAttention easy to adopt: no quality trade-off to argue about.
The ratio is 4x here because the tile is a quarter of the sequence. The real saving grows with sequence length, because the tile stays a fixed size while N grows. Memory goes from growing with the square of the length to growing with the length.
The last three lines are the reason anyone cares. At 131,072 tokens, one attention head's score matrix is 34 GB. A model has dozens of heads and dozens of layers. Without tiling, long context is not expensive, it is impossible.
The three running quantities
m is the largest score seen so far in each row. Subtracting it before exp stops the exponential from overflowing. This is the standard safe-softmax trick applied to a stream.
l is the running sum of exponentials, which is the softmax denominator.
O is the running weighted sum of value vectors.
correction = np.exp(m - m_new) is the line that earns the whole algorithm. When a new tile raises the row maximum, everything accumulated under the old maximum was scaled by the wrong constant. Multiplying by this factor repairs l and O together, exactly.
Measuring it on a real GPU
PyTorch hides several attention kernels behind one function. MATH is the textbook version that builds the score matrix. The tiled kernels do not.
import torch, torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
assert torch.cuda.is_available(), "this measurement needs a CUDA GPU"
print("gpu:", torch.cuda.get_device_name(0), "| torch:", torch.__version__)
B, H, N, D = 1, 16, 8192, 64
q, k, v = (torch.randn(B, H, N, D, device="cuda", dtype=torch.float16) for _ in range(3))
def works(backend):
try:
with sdpa_kernel(backend):
F.scaled_dot_product_attention(q[:, :, :64], k[:, :, :64], v[:, :, :64])
return True
except RuntimeError:
return False
for b in (SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION,
SDPBackend.CUDNN_ATTENTION, SDPBackend.MATH):
print(f" backend {b.name:20s} usable here: {works(b)}")
def peak_mb(backend):
torch.cuda.synchronize(); torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
base = torch.cuda.memory_allocated()
with sdpa_kernel(backend):
out = F.scaled_dot_product_attention(q, k, v)
torch.cuda.synchronize()
mb = (torch.cuda.max_memory_allocated() - base) / 2**20
return out, mb
o_math, mb_math = peak_mb(SDPBackend.MATH)
o_tiled, mb_tiled = peak_mb(SDPBackend.EFFICIENT_ATTENTION)
print()
print("same answer:", torch.allclose(o_math, o_tiled, atol=2e-3))
print(f"MATH backend peak extra memory: {mb_math:8.1f} MB")
print(f"tiled backend peak extra memory: {mb_tiled:8.1f} MB")
print(f"the N x N scores alone would be {B*H*N*N*2/2**20:8.1f} MB in fp16")gpu: NVIDIA RTX A6000 | torch: 2.5.1+cu121 backend FLASH_ATTENTION usable here: False backend EFFICIENT_ATTENTION usable here: True backend CUDNN_ATTENTION usable here: True backend MATH usable here: True same answer: True MATH backend peak extra memory: 13440.1 MB tiled backend peak extra memory: 16.0 MB the N x N scores alone would be 2048.0 MB in fp16
PyTorch also prints several UserWarning lines to stderr explaining why each rejected kernel was rejected. Read them, they are unusually informative.
These numbers come from one machine: an NVIDIA RTX A6000 running Windows. Peak memory is deterministic for a fixed shape, so the same shapes reproduce the same figures. The list of usable backends will differ on your machine, which is exactly why the script prints it.
13,440 MB against 16 MB. The MATH path is worse than the raw 2,048 MB of scores because it promotes to float32 and holds several intermediates at once. The tiled path allocates the output and little else.
FLASH_ATTENTION reports False here. This Windows PyTorch wheel was not built with that kernel, and the card is Ampere rather than Hopper. The memory-efficient kernel uses the same tiling idea, so the lesson holds. Print the table rather than assuming: the answer depends on your operating system, GPU generation, dtype and head dimension.
Common mistakes
Expecting a speed-up on short sequences. At 128 or 512 tokens the score matrix fits in cache anyway. Tiled kernels win when the sequence is long, and can lose below a few hundred tokens.
Passing float32. The fastest kernels want float16 or bfloat16. Hand them float32 and PyTorch falls back to a slower path without complaining. Print the backend you actually got.
Passing a dense mask when is_causal=True would do. A boolean mask of shape N x N re-creates the exact allocation you were trying to avoid. Use the flag.
Installing flash-attn and assuming it is in use. pip install flash-attn --no-build-isolation needs a matching CUDA toolkit and a long compile. Check the backend per call, not the install log.
Reading a speed claim without its hardware. FlashAttention-2 targets Ampere through Ada. FlashAttention-3 is Hopper only. FlashAttention-4 targets Hopper and Blackwell. A benchmark from one generation says nothing about another.
Try it yourself
Change B in the NumPy script to 512, so there is one tile. Confirm the memory ratio becomes 1x and the answer does not move. Then set it to 16 and watch the correction step fire far more often. Print m inside the loop to see the row maximum climb.
What to learn next
- Memory-bound vs compute-bound — why moving data, not multiplying, sets the clock.
- Attention — the operation FlashAttention makes cheap, from first principles.
- Context window — what a longer context buys you, and what it costs.
Researcher — Mathematics and papers.
The problem being solved
Standard attention for one head, with $Q, K, V \in \mathbb{R}^{N \times d}$:
$$ O = \operatorname{softmax}!\left(\frac{QK^{\top}}{\sqrt{d}}\right) V $$
$N$ is sequence length, $d$ is head dimension, and $\sqrt{d}$ keeps the logits' variance stable at initialisation. The intermediates $S = QK^{\top}/\sqrt{d}$ and $P = \operatorname{softmax}(S)$ are each $N \times N$.
FLOPs are $\Theta(N^2 d)$ and cannot be reduced without changing the function. Memory traffic can be. Writing and re-reading $S$ and $P$ costs $\Theta(N^2)$ accesses to high-bandwidth memory (HBM), and on modern accelerators that traffic, not the arithmetic, sets wall-clock time.
Online softmax
The recurrence comes from Milakov and Gimelshein (2018), Online normalizer calculation for softmax (arXiv:1805.02867). Process key blocks $j = 1 \dots T$. After block $j$, hold $m^{(j)}$, the running row maximum, $\ell^{(j)}$, the running exponential sum, and $O^{(j)}$, the running unnormalised output.
$$ m^{(j)} = \max!\left(m^{(j-1)},\; \operatorname{rowmax}(S^{(j)})\right) $$
$$ \ell^{(j)} = e^{m^{(j-1)} - m^{(j)}} \ell^{(j-1)} + \operatorname{rowsum}!\left(e^{S^{(j)} - m^{(j)}}\right) $$
$$ O^{(j)} = e^{m^{(j-1)} - m^{(j)}} O^{(j-1)} + e^{S^{(j)} - m^{(j)}} V^{(j)} $$
$S^{(j)} \in \mathbb{R}^{N \times B_c}$ is the score block for key tile $j$ of width $B_c$, and $V^{(j)}$ is the matching value tile. The output is $O^{(T)} / \ell^{(T)}$. The factor $e^{m^{(j-1)} - m^{(j)}}$ rescales prior partial sums onto the new maximum, so the result matches the one-shot computation up to floating-point reassociation.
IO complexity
Dao, Fu, Ermon, Rudra and Ré (2022), FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (arXiv:2205.14135, NeurIPS 2022) prove HBM accesses of
$$ \Theta!\left(\frac{N^2 d^2}{M}\right) $$
against $\Theta(Nd + N^2)$ for the standard implementation, where $M$ is on-chip SRAM size in elements. For $d = 64$ and $M$ of order $10^5$ this is a 5–20x traffic reduction. They further show it is IO-optimal: no exact attention algorithm uses asymptotically fewer HBM accesses for all $M$.
Activation memory falls from $\Theta(N^2)$ to $\Theta(N)$, because $S$ and $P$ are never stored. The backward pass recomputes them tile by tile from the saved $\ell$ and $m$ statistics. That recomputation is close to free in wall-clock terms precisely because the kernel is memory-bound.
Rabe and Staats (2021), Self-attention Does Not Need $O(n^2)$ Memory (arXiv:2112.05682) reached the low-memory result independently. FlashAttention's contribution is the fused CUDA kernel that makes it faster rather than only smaller.
Version lineage
| Version | Year | Target | Key change |
|---|---|---|---|
| FlashAttention | 2022 | Ampere | Tiling plus recomputation; IO-optimal exact attention |
| FlashAttention-2 | 2023 | Ampere, Ada, Hopper | Fewer non-matmul FLOPs, parallelism over sequence, better warp partitioning; about 2x over v1 |
| FlashAttention-3 | 2024 | Hopper only | Warp-specialised producer/consumer pipelining, FP8 forward |
| FlashAttention-4 | 2026 | Hopper, Blackwell | CuTeDSL rewrite, fully asynchronous MMA, larger tiles, software-emulated exponentials |
- Dao (2023), FlashAttention-2 — arXiv:2307.08691, ICLR 2024.
- Shah, Bikshandi, Zhang, Thakkar, Ramani and Dao (2024), FlashAttention-3 — arXiv:2407.08608.
- Zadouri, Hoehnerbach, Shah, Liu, Thakkar and Dao (2026), FlashAttention-4 — arXiv:2603.05451, reporting roughly 1613 TFLOP/s on B200.
The version numbers track GPU generations, not algorithmic depth. The mathematics has not moved since 2022. Each release re-targets a new memory hierarchy and a new asynchronous instruction set, so read every FlashAttention benchmark as a statement about one chip.
What FlashAttention does not fix
Compute still scales as $\Theta(N^2 d)$. At long enough $N$ arithmetic dominates again, and only sparsity or a different operator helps — see sparse and block attention patterns and linear attention.
The KV cache is untouched. During decoding, memory is dominated by cached keys and values, which is the province of multi-query attention and multi-head latent attention.
Numerics are equivalent up to reassociation, not bitwise identical. Tests asserting exact equality against a reference implementation will fail, and should be written with a tolerance.
Beyond hand-written kernels
torch.nn.attention.flex_attention compiles a user-supplied score_mod or mask_mod into a fused tiled kernel, so a new masking scheme does not require new CUDA. It is documented as a prototype API as of PyTorch 2.13, so pin your version and expect signature churn. The FlashAttention-4 authors report that their CuTeDSL framework hosts FlexAttention-style and block-sparse variants without changes to the core. That is the direction of travel: one tiled skeleton, many masks.
What to learn next
- Memory-bound vs compute-bound — why moving data, not multiplying, sets the clock.
- Attention — the operation FlashAttention makes cheap, from first principles.
- Context window — what a longer context buys you, and what it costs.