Fast attention with scaled_dot_product_attention
F.scaled_dot_product_attention computes exactly what the textbook formula computes, but through fused kernels like FlashAttention that never write the huge score matrix to memory.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
scaled_dot_product_attention gives the same answer as hand-written attention, while skipping the giant table that makes attention expensive.
Imagine seating guests at a wedding. One way to do it: draw a table scoring every guest against every other guest, all pairs, on paper. Then read it to decide the seating. For two hundred guests that table has forty thousand cells. The paper is the problem, not the thinking.
A clever planner scores guests section by section, keeps a small running tally, and never writes the full table anywhere.
Why it exists
Attention — the mechanism that lets a transformer decide which words matter to which — compares every position with every other position. Double the text length, and the comparison table grows four times. For long documents, that table becomes the single biggest thing in GPU memory.
The insight behind FlashAttention: the table never needs to exist all at once. Computing it in small tiles, keeping running totals, gives the identical answer while touching memory far less. Reading and writing memory, not arithmetic, is what GPUs wait for — so skipping the table is also faster.
PyTorch wraps this behind one function, and picks the best available method for your hardware by itself.
How it works
by hand: compare all pairs -> giant score table in memory -> answer
(grows 4x when text doubles)
fused: compare tile by tile -> running tally -> same answer
(giant table never exists)A real example you have seen
Chat assistants accepting whole documents — long chats, long PDFs — became normal only after tricks like this. With the full-table method, memory for a hundred-page context is unaffordable. The tiled method is one reason "paste your whole file" works today.
Remember this
- Attention compares everything with everything; the score table dwarfs everything else at long lengths.
- The fused method computes the same answer without materialising the table.
- In PyTorch, this is one function call, and it chooses the fastest method itself.
What to learn next
- Mixed precision training — the dtype prerequisite for the fast backends.
- Transformers — the architecture this function accelerates.
- torch.compile — fusing everything around the attention call.
Developer — Code and libraries.
Setup
pip install torchCorrectness runs anywhere, CPU included. The memory comparison needs a GPU (captured with torch 2.5.1 on an NVIDIA RTX A6000).
Same answer as the textbook formula
import torch
import torch.nn.functional as F
torch.manual_seed(0)
q = torch.randn(2, 4, 128, 64) # batch, heads, sequence, head_dim
k = torch.randn(2, 4, 128, 64)
v = torch.randn(2, 4, 128, 64)
# attention written by hand, straight from the paper
scores = q @ k.transpose(-2, -1) / (64 ** 0.5)
by_hand = torch.softmax(scores, dim=-1) @ v
fused = F.scaled_dot_product_attention(q, k, v)
print("same answer:", torch.allclose(by_hand, fused, atol=1e-5))
print("output shape:", tuple(fused.shape))same answer: True output shape: (2, 4, 128, 64)
Drop-in: if your module contains the three-line hand version, replacing it with the one-liner changes nothing about the result. Causal masking is the is_causal=True argument; arbitrary masks go in attn_mask.
The memory difference, measured
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
if not torch.cuda.is_available():
raise SystemExit("needs a GPU to show the memory difference")
def peak_mb(backend):
torch.manual_seed(0)
q = torch.randn(1, 8, 4096, 64, device="cuda", dtype=torch.float16)
k, v = torch.randn_like(q), torch.randn_like(q)
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
with sdpa_kernel([backend]):
F.scaled_dot_product_attention(q, k, v)
torch.cuda.synchronize()
return torch.cuda.max_memory_allocated() / 1024**2
print(f"math backend (stores the score matrix): {peak_mb(SDPBackend.MATH):8.1f} MB peak")
# FlashAttention first; some builds (Windows wheels) ship without that kernel,
# so fall back to the memory-efficient one, which skips the score matrix too.
for label, backend in [("flash", SDPBackend.FLASH_ATTENTION),
("mem-efficient", SDPBackend.EFFICIENT_ATTENTION)]:
try:
print(f"{label:<13} backend (never stores it): {peak_mb(backend):8.1f} MB peak")
break
except RuntimeError:
continue # kernel absent in this buildmath backend (stores the score matrix): 1716.2 MB peak mem-efficient backend (never stores it): 24.1 MB peak
Seventy times less memory, at sequence length 4096. The score matrix alone is 8 heads x 4096 x 4096 in half precision — 256 MB that the math path must write, read and keep for backward, while the fused kernels never build it at all.
The loop exists because the machine that captured this output runs a Windows
build with no FlashAttention kernel compiled in, so SDPBackend.FLASH_ATTENTION
raises RuntimeError: No available kernel. The memory-efficient kernel uses the
same tiling idea and lands on the same number. On a Linux build the first branch
wins and the line reads flash. Print rather than assume: which kernels exist
depends on your operating system, GPU generation, dtype and head dimension. The
FlashAttention lesson shows how to
list the available backends directly.
The sdpa_kernel context restricts which backend may run — here used to stage the comparison. In normal code you write no context at all and PyTorch dispatches: FlashAttention when dtype and hardware allow, a memory-efficient variant otherwise, and the explicit math fallback last.
Backend requirements matter. FlashAttention needs float16 or bfloat16 on a reasonably modern NVIDIA card, and head dimensions within kernel limits (typically ≤ 128 covers common models). Feed float32 and you silently get a slower backend — pair this with mixed precision.
Common mistakes
Keeping the hand-written version out of habit. Every hand-rolled softmax-attention forgoes the fused kernels and stores the score matrix for backward. One function call is the whole migration for most modules.
Wrong tensor layout. The function expects (batch, heads, seq, head_dim). Code holding (batch, seq, heads, head_dim) needs a transpose(1, 2) first — and shape errors here produce wrong attention, not always a crash. See reshape, view and contiguity.
Benchmarking it in float32. The float32 path cannot use FlashAttention, so the comparison shows "no speedup" and the wrong conclusion gets drawn. Benchmark in half precision, with honest timing.
Assuming the answer is bit-identical. Tiled accumulation reorders floating-point sums; results match to tolerance (allclose), not exactly. Tests using torch.equal will fail for no real reason.
Try it yourself
In sdpa_mem.py, double the sequence length to 8192 and predict both numbers before running: the math backend should roughly quadruple, the flash backend roughly double. Then add is_causal=True to both and see whether the gap survives masking.
What to learn next
- Mixed precision training — the dtype prerequisite for the fast backends.
- Transformers — the architecture this function accelerates.
- torch.compile — fusing everything around the attention call.
Researcher — Mathematics and papers.
The IO-aware argument
Standard attention on sequence length $N$, head dimension $d$:
$$ S = \frac{QK^\top}{\sqrt{d}}, \quad P = \mathrm{softmax}(S), \quad O = PV $$
- $Q, K, V \in \mathbb{R}^{N \times d}$ — queries, keys, values (per head).
- $S, P \in \mathbb{R}^{N \times N}$ — scores and attention weights: the $O(N^2)$ objects.
FLOPs are $O(N^2 d)$ in every variant; FlashAttention's contribution is HBM traffic. Materialising $S$ and $P$ costs $O(N^2)$ reads/writes to device memory. Tiling $Q$ into row blocks and $K, V$ into column blocks sized to on-chip SRAM, and maintaining the online softmax running statistics — row maximum $m_i$ and normaliser $\ell_i$, rescaling partial outputs as new tiles arrive — yields exact attention with $O(N^2 d^2 / M)$ HBM accesses for SRAM size $M$. Attention is memory-bound at these intensities, so the traffic reduction converts directly to wall-clock speedup (2–4x typical) and the working set drops from $O(N^2)$ to $O(N)$: the numbers in the developer demo.
The backward pass recomputes $P$ tile-by-tile from stored $(O, m, \ell)$ rather than reading a saved $N^2$ tensor — gradient checkpointing's idea, specialised and made free.
Lineage and variants
- Milakov and Gimelshein (2018), Online normalizer calculation for softmax — the running-softmax primitive.
- Rabe and Staats (2021), Self-attention Does Not Need $O(n^2)$ Memory — the memory-efficient backend's ancestor.
- Dao et al. (2022), FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, NeurIPS.
- Dao (2023), FlashAttention-2 — work partitioning across warps, ~2x again; FlashAttention-3 (Shah et al., 2024) targets Hopper's TMA and FP8.
Contrast with approximate efficient attentions (Performer, Linformer, sliding windows): SDPA's backends are exact. The approximation family trades quality for asymptotics; FlashAttention changed the constant factors so thoroughly that exact attention remained the default at production context lengths. Inference-side descendants (PagedAttention/vLLM) apply the same IO lens to KV-cache management.
torch.compile integration: SDPA is a single graph node the compiler pattern-matches (fusing surrounding projections); hand-written attention decomposes into ops it must fuse piecemeal — one more reason the one-liner wins.
References
All above, plus the PyTorch SDPA documentation for the dispatch rules and per-backend constraints (dtype, head_dim, mask support), which shift by release.
What to learn next
- Mixed precision training — the dtype prerequisite for the fast backends.
- Transformers — the architecture this function accelerates.
- torch.compile — fusing everything around the attention call.