Attention Is the Bottleneck: FlashAttention, MLA, Sparse and Linear Attention in 2026
Every transformer layer has one operation that gets more expensive as the context grows: attention. The FFN costs the same per token whether your prompt is 100 tokens or 100,000. Attention doesn't. Most of the interesting inference work of the last four years has been some version of "make attention cheaper without making the model dumber."
I've touched pieces of this in older posts: PagedAttention in vLLM, KV cache compression on my MacBook, and LFM2 vs Qwen3.5, where both models throw out most of their attention layers. This post puts it all on one map.
TL;DR
- Attention is slow mostly because of memory traffic, not FLOPs. Prefill wastes bandwidth writing the N x N score matrix out to HBM. Decode has to read the entire KV cache for every token it generates.
- FlashAttention fixes the prefill side by never materializing that score matrix. Output is exact, the attention matrix itself takes no memory, and each version since v1 has been rewritten around a newer NVIDIA GPU (A100, H100, B200).
- GQA and MLA fix the decode side by making the KV cache smaller. Llama 3 8B with GQA stores 128 KiB per token. The same model with full multi-head attention would store 512 KiB.
- Sparse attention (sliding windows, DeepSeek's NSA and DSA) and linear attention hybrids (Gated DeltaNet, Kimi Delta Attention, LFM2's short convolutions) stop every token from attending to every other token.
- In practice: use the fused kernel your framework already ships, pick models with small KV caches, and let your serving engine handle paging and prefix caching. Writing your own attention kernel is almost never the right move.
What attention costs
The formula fits on one line:
Attention(Q, K, V) = softmax(Q Kᵀ / √d) V
For a sequence of N tokens with head dimension d, Q Kᵀ is an N x N matrix of scores. At 32K tokens that is a billion entries per head per layer. The FLOPs grow as N² too, but GPUs have a lot of FLOPs. The bigger problem is where that matrix lives.
A GPU has two kinds of memory that matter here. HBM is big and slow (80 GB at about 3.35 TB/s on an H100). On-chip SRAM is tiny and fast (a couple hundred KB per streaming multiprocessor, roughly an order of magnitude more bandwidth). The textbook implementation runs as three separate kernels:
- Compute
S = Q Kᵀ, write S to HBM - Read S, compute
P = softmax(S), write P to HBM - Read P and V, compute
O = P V, write O to HBM
Steps 1 and 2 shuffle an N x N matrix through slow memory twice. The tensor cores spend most of their time waiting for data. That's what the FlashAttention paper meant by calling attention IO-bound.
Prefill and decode are different problems
Inference has two phases, and attention hurts in each one for a different reason.
Prefill reads your prompt. All N query tokens get processed at once, so there's real matrix multiplication to do and the GPU can stay busy. The waste is the N x N intermediate.
Decode generates one token at a time. There's a single query vector, but it still has to attend to every cached key and value. You do very little math per byte loaded, so decode is bound purely by memory bandwidth. A quick calculation for Llama 3 8B (32 layers, 8 KV heads, head dim 128, BF16):
KV bytes per token = 2 (K and V) x 32 layers x 8 heads x 128 dims x 2 bytes
= 131,072 bytes = 128 KiB
At 32K context: 32,768 x 128 KiB = 4 GiB read per generated token, per sequence
On an H100 (3.35 TB/s): ~1.3 ms of pure KV reads per token
That 1.3 ms gets paid by every sequence in the batch on every decode step, and it's before any weights are loaded. So the work splits into two camps. Kernel people (FlashAttention) fix the IO pattern, which mostly helps prefill. Architecture people (GQA, MLA, sparse, linear) shrink what decode has to read.
FlashAttention: same math, different IO
FlashAttention (Tri Dao et al., 2022) runs all three steps as one fused kernel and never writes S or P to HBM at all. It loads a tile of Q into SRAM, streams tiles of K and V past it, and accumulates the output on-chip.
The catch is softmax. To normalize a row you need its maximum and its sum over every key, and you only ever see one tile at a time. The fix is online softmax (Milakov and Gimelshein, 2018): keep a running max and a running denominator, and every time a new tile raises the max, rescale what you've accumulated so far. Here's that idea in numpy:
import numpy as np
def naive_attention(q, K, V):
s = K @ q / np.sqrt(q.shape[0]) # all N scores at once
p = np.exp(s - s.max())
return (p / p.sum()) @ V
def tiled_attention(q, K, V, block=64):
d = q.shape[0]
m, l = -np.inf, 0.0 # running max, running denominator
acc = np.zeros(V.shape[1]) # running (unnormalized) output
for i in range(0, K.shape[0], block):
s = K[i:i+block] @ q / np.sqrt(d) # one tile of scores
m_new = max(m, s.max())
scale = np.exp(m - m_new) # fix up everything seen so far
p = np.exp(s - m_new)
l = l * scale + p.sum()
acc = acc * scale + p @ V[i:i+block]
m = m_new
return acc / l
rng = np.random.default_rng(0)
q = rng.standard_normal(128)
K = rng.standard_normal((4096, 128))
V = rng.standard_normal((4096, 128))
print(np.abs(naive_attention(q, K, V) - tiled_attention(q, K, V)).max())
# 1.39e-16The two outputs differ by 1.4e-16, which is floating-point noise. FlashAttention is exact, not an approximation. That's a big part of why it took over: you got the speed and didn't have to argue with anyone about accuracy.
For the backward pass, FlashAttention doesn't store the attention matrix. It recomputes it from Q, K, V and the saved softmax statistics. That costs more FLOPs, but FLOPs were cheap and memory traffic wasn't, so it ends up faster anyway. The paper reported 3x faster GPT-2 training and attention memory that grows linearly with sequence length instead of quadratically. That second result is a big reason context windows jumped from 2K to 32K and beyond.
FlashAttention-2: better use of the GPU
FlashAttention-2 (2023) kept the algorithm and fixed how the work was split across the GPU:
- Fewer non-matmul FLOPs. Tensor cores are much faster at matmuls than the regular units are at everything else, so FA2 delays the final normalization and trims the bookkeeping.
- Parallelism over sequence length, not just batch and heads. With long sequences and small batches, FA1 left SMs sitting idle.
- Better split inside a block. Each warp takes a slice of Q instead of a slice of K, so warps don't have to sync through shared memory to combine partial results.
The result was about 2x over FA1, reaching 50 to 73% of the A100's theoretical peak. That's close to what a plain matmul gets.
Flash-Decoding: parallelism for decode
FA2 parallelizes over queries, and decode has exactly one. With a small batch and a long context, most of the GPU had nothing to do. Flash-Decoding (late 2023) splits the KV cache into chunks along the sequence, attends to each chunk in parallel, and combines the partial results with the same log-sum-exp trick from online softmax. On very long sequences it reported up to 8x faster decoding. Some version of split-KV is now standard in every serious decode kernel.
FlashAttention-3: built for Hopper
The H100 added hardware FA2 couldn't use: the Tensor Memory Accelerator (TMA) for async copies, warpgroup-wide matrix instructions, and FP8. FlashAttention-3 (2024) was written for it:
- Warp specialization. Some warps only load data and others only compute, so loads and math overlap.
- Ping-pong scheduling. Softmax uses the slow exponential unit while the matmuls use tensor cores. FA3 interleaves two warpgroups so one runs softmax while the other runs its GEMM.
- FP8 with incoherent processing. It multiplies Q and K by a random orthogonal (Hadamard) matrix before quantizing, which spreads outliers across dimensions. Numerical error came out 2.6x lower than a baseline FP8 attention.
FP16 reached up to 740 TFLOPs/s (about 75% of H100 peak), 1.5 to 2x faster than FA2. FP8 got close to 1.2 PFLOPs/s.
FlashAttention-4: built for Blackwell
The B200 roughly doubled tensor core throughput, but the exponential unit and shared memory bandwidth didn't keep pace. Matmul isn't the bottleneck anymore. Softmax is. FlashAttention-4 (MLSys 2026) is built around that imbalance:
- Software exponentials. Part of the
expwork moves off the dedicated special-function unit and gets computed as a polynomial on the regular FMA units, so both run at once. - Conditional rescaling. Online softmax normally rescales the accumulator every time the running max goes up. FA4 skips the rescale unless the max moved enough to threaten numerical stability, which cuts rescaling about 10x.
- Tensor memory and 2-CTA MMA to take pressure off shared memory, especially in the backward pass.
- Written in CuTe-DSL (Python) instead of C++ templates, and it compiles 20 to 30x faster.
On a B200 in BF16 it reports up to 1,613 TFLOPs/s (71% utilization), 1.3x faster than cuDNN 9.13 and 2.7x faster than Triton.
| Version | Year | Target GPU | Main idea | Headline |
|---|---|---|---|---|
| FlashAttention | 2022 | A100 | Tiling + online softmax, recompute in backward | Linear memory, ~3x GPT-2 training |
| FlashAttention-2 | 2023 | A100 | Fewer non-matmul FLOPs, sequence parallelism | ~2x over FA1, 50-73% of peak |
| Flash-Decoding | 2023 | Any | Split the KV cache for decode | Up to 8x on long-context decode |
| FlashAttention-3 | 2024 | H100 | Async warp specialization, ping-pong, FP8 | 740 TFLOPs/s FP16 |
| FlashAttention-4 | 2026 | B200 | Software exp, conditional rescaling | 1,613 TFLOPs/s BF16 |
Each new GPU generation made matmuls faster than everything around them, and each new FlashAttention rewrite moved work off whatever had become the slowest unit. Four papers, one habit: figure out which unit is actually saturated before you optimize anything.
Shrinking the KV cache
FlashAttention doesn't change how many bytes decode has to read. The only way to read fewer bytes is to store fewer, and that's an architecture choice made before training.
MQA and GQA
Standard multi-head attention (MHA) gives each query head its own K and V heads. Multi-Query Attention (Shazeer, 2019) has every query head share one K/V head. That shrinks the cache by the head count but costs some quality. Grouped-Query Attention (Ainslie et al., 2023) is the middle ground: query heads are split into groups and each group shares a K/V head. Llama 3 uses 32 query heads and 8 KV heads, which is a 4x smaller cache for almost no quality loss. GQA is now the default for dense models.
MLA
DeepSeek-V2 introduced Multi-head Latent Attention, and it goes further. Instead of caching K and V, it caches one small latent vector per token per layer and projects it back up to per-head keys and values at attention time. The up-projection can be folded into the query and output projections, so you never have to materialize the full K and V during decode. DeepSeek reported a 93.3% smaller KV cache and 5.76x higher max generation throughput than their earlier 67B dense model.
These per-token numbers are BF16, computed from each model's config:
| Model | Attention | KV per token | KV at 128K context |
|---|---|---|---|
| Llama 3 8B, if it used MHA | 32 KV heads | 512 KiB | 64 GiB |
| Llama 3 8B | GQA, 8 KV heads | 128 KiB | 16 GiB |
| Llama 3 70B | GQA, 8 KV heads | 320 KiB | 40 GiB |
| DeepSeek-V3 (671B) | MLA, 512 + 64 dims | 68.6 KiB | 8.6 GiB |
DeepSeek-V3 has almost 10x the total parameters of Llama 3 70B and stores about a fifth of the KV cache per token. That one still surprises me every time I do the math. MLA has since shown up in Kimi K2 and a growing list of other models.
Quantize what's left
You can also store the cache in fewer bits. FP8 KV cache is a routine flag in vLLM and SGLang now, and 4-bit schemes work if you handle outliers carefully. I benchmarked TurboQuant+ on Apple Silicon earlier this year. The speedups were real, but quality fell off a cliff at the most aggressive settings, so measure on your own workload before shipping it.
Paging and prefix sharing
The last memory fix is about allocation, not size. PagedAttention stores the cache in fixed-size blocks instead of one contiguous buffer per request, so memory isn't wasted on padding. SGLang's RadixAttention keeps a prefix tree of cached blocks, so requests that start with the same system prompt or few-shot examples reuse the same KV. That's the same mechanism behind provider-side prompt caching.
Attending to fewer tokens
Everything so far keeps full attention, where every token can see every earlier token. The next family drops that.
Sliding windows and local:global layers
The simplest version is to let each token see only the last W tokens. Mistral 7B used a 4,096-token window. The obvious problem is that the model can't directly look at anything older than the window. The fix that stuck is interleaving: most layers are local and a few are global. Gemma 3 runs 5 local layers (1,024-token window) for every global layer, which keeps the KV cache for most layers tiny. gpt-oss alternates banded-window and full layers too.
A related finding is attention sinks. StreamingLLM noticed that models dump a lot of attention onto the very first tokens no matter what they contain, and that evicting those tokens breaks generation. Keep a handful of sink tokens plus a recent window, and the model streams much longer than its training length. gpt-oss bakes this in with a learned sink per head.
Learned sparsity: NSA and DSA
Fixed windows decide in advance which tokens matter. DeepSeek's two papers learn which ones matter instead.
Native Sparse Attention (NSA, 2025) runs three branches per query: attention over compressed block summaries for the big picture, attention over the top-scoring full-resolution blocks for detail, and a sliding window for local context. A learned gate mixes them, and the blocks are sized so the kernel stays fast on real hardware. It's trained this way from the start, which is where "native" comes from.
DeepSeek Sparse Attention (DSA), introduced in V3.2, is simpler. A small, cheap lightning indexer scores every previous token for the current query. Only the top 2,048 go into real attention. The indexer itself is still quadratic, but it uses few heads and runs in FP8, so it's cheap next to the main attention, which is now O(N x k) instead of O(N²). DeepSeek cut its API prices by more than half when it shipped.
Replacing attention in most layers
The most aggressive option is to stop using softmax attention in most layers. Linear attention variants replace the growing KV cache with a fixed-size recurrent state. That means constant memory and constant time per token, no matter how long the context is. Pure linear models have historically been worse at exact recall, so everyone shipping them in 2025 and 2026 runs a hybrid that keeps a few full attention layers for that.
- Qwen3-Next and Qwen3.5 use Gated DeltaNet in three of every four layers.
- Kimi Linear uses Kimi Delta Attention (KDA, finer-grained gating than Gated DeltaNet) at 3 KDA : 1 MLA. Moonshot reports 75% less KV cache and up to 6.3x faster per-token decode at 1M context.
- LFM2 skips linear attention altogether and uses gated short convolutions (kernel length 3) in 10 of its 16 layers, with GQA in the other 6.
I measured that last pair against each other in LFM2 vs Qwen3.5 on an M4. At equal parameter count, LFM2 decoded 1.2 to 1.6x faster. The model cards don't advertise the downside: hybrids need custom kernels, and support in llama.cpp, MLX and vLLM tends to lag behind new architectures by weeks or months.
What I'd actually do
- Don't write an attention kernel.
torch.nn.functional.scaled_dot_product_attentionalready dispatches to FlashAttention, cuDNN or memory-efficient backends. vLLM and SGLang ship FlashAttention and FlashInfer. MLX has a fusedscaled_dot_product_attention. If you're calling a naivesoftmax(q @ k.T) @ vanywhere in production, fixing that is your quick win. - Choose models by KV cache size, not just parameter count. For long-context or high-concurrency serving, an MLA or hybrid model can fit several times more concurrent sequences in the same memory as a GQA model of similar size.
- Let the serving engine handle memory. Paging, prefix caching, and FP8 KV cache are config flags now. Turn them on before you consider anything exotic.
- On edge devices, the architecture matters more than the kernel. A phone or laptop has a fraction of a datacenter GPU's memory bandwidth, so the decode cost from earlier is everything. Hybrids like LFM2 and Qwen3.5, plus a quantized cache, are where the real wins are.
- Measure the slow resource. FlashAttention-1 through 4 all came from asking what the hardware was actually waiting on. Do the same with your own workload before you pick a fix.
Nine years after "Attention Is All You Need," attention is still the expensive part. The difference is that we now know exactly which bytes and which execution units make it expensive, and every technique in this post goes after one of them.
Related posts
LFM2 vs Qwen3.5 on an Apple M4: Every Small Model, Both Backends
llama-bench plus perplexity plus a six-prompt smell test across all three LFM2 sizes (350M, 700M, 1.2B) and the two smallest Qwen3.5 models (0.8B, 2B) on my M4, on Metal and on CPU alone. Liquid's hybrid does win at decode, but the quantization result surprised me more.
TurboQuant+ Meets Gemma on a Modal L40S
Second pass on TurboQuant+ KV cache compression, this time on a rented L40S across Gemma 3 12B, Gemma 4 E4B, and Gemma 4 26B MoE. One works beautifully, two break in interesting ways.
Benchmarking TurboQuant+ KV Cache Compression on Apple Silicon
I tested TurboQuant+ KV cache compression across 1.5B, 7B, and 14B models on an M4 MacBook Air. The speed gains are real, but there are sharp cliffs you need to know about.