Fast Attention and Long Context
Memory-bound vs compute-bound
Most of what looks slow in a language model is not arithmetic, it is waiting for numbers to arrive from memory, and knowing which one you are in tells you what to fix.
- 12 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
A program is either waiting for numbers to arrive, or busy doing sums with them. Almost every slow language model is waiting.
The restaurant you have eaten in
Picture a busy restaurant with twenty quick cooks. All the vegetables live in a store room, reached through one narrow doorway.
The cooks are never the problem. They stand around while one person squeezes through the doorway with a crate. Hire ten more cooks and the food comes out at the same speed.
Widen the doorway, and everything changes at once. That doorway is what engineers call memory bandwidth. It is how fast data travels from storage into the part that does the work.
The two situations
A piece of work is memory-bound when the doorway is the limit. The processing units sit idle waiting for data.
A piece of work is compute-bound when the doorway is wide enough and the arithmetic itself is the limit.
memory-bound compute-bound
------------ -------------
[ store room ] --narrow--> [ store room ] ==wide==>
cooks: idle cooks: flat out
fix: move less data fix: do fewer sums,
or reuse what you moved or use faster hardwareGetting this backwards wastes months. A team buys a faster chip to fix a doorway problem, and nothing improves.
Why language models are usually waiting
Writing one word means reading every weight in the model. A weight is one of the stored numbers the model learned.
For a large model that is tens of gigabytes, read from memory, to produce a single word. Then read all of them again for the next word.
The amount of arithmetic per weight is tiny. Each weight is fetched, used for one multiply, and dropped. That is the worst possible ratio.
The one trick that fixes it
Serve many users at once. Fetch each weight once, and use it for fifty requests instead of one.
The crate comes through the doorway once and feeds fifty plates. The doorway stops being the limit. This is called batching. It is why a hosted model costs far less per word than the same model answering you alone.
That is also why a chatbot answering only you feels wasteful of an expensive graphics card. It is. The card is mostly idle, waiting.
A useful number to remember
Ask: for every byte I move, how many sums do I do?
Low means waiting. High means working. Engineers call this ratio arithmetic intensity, and it decides which of the two worlds you are in.
Where you have already seen this
- Your phone's camera app running an effect instantly, while a large model on the same phone crawls.
- A cloud model that gets cheaper per word the busier the service is.
- A file copy that pins the disk at full speed while the processor idles.
Remember this
- Memory-bound means waiting for data. Compute-bound means busy with sums.
- Writing words one at a time is memory-bound, almost always.
- Batching many requests together is the standard cure, because each weight gets reused.
What to learn next
- Multi-query attention — the first serious attack on decode-time memory traffic.
- Latency and throughput — the two numbers this trade-off is measured in.
- Is the GPU waiting for data? — the same question, one level up the stack.
Developer — Code and libraries.
Setup
pip install torch # the second script needs a CUDA GPU; the first needs nothingWritten against PyTorch 2.5.1+cu121, Python 3.10.
Step one: count, do not time
Timing is noisy and machine-specific. Counting is neither. For a matrix multiply of shape (B, K) @ (K, N) the arithmetic is fixed and so is the traffic.
def intensity(M, K, N, bytes_per_el=2):
flops = 2 * M * K * N # one multiply + one add per element
moved = (M * K + K * N + M * N) * bytes_per_el # inputs in, output out
return flops, moved, flops / moved
print("one 4096x4096 weight matrix, bf16, applied to a batch of B tokens")
print(f"{'B':>6} {'FLOPs':>15} {'bytes moved':>15} {'FLOPs/byte':>12}")
for B in (1, 8, 64, 512, 4096):
f, m, ai = intensity(B, 4096, 4096)
print(f"{B:>6} {f:>15,} {m:>15,} {ai:>12.2f}")one 4096x4096 weight matrix, bf16, applied to a batch of B tokens
B FLOPs bytes moved FLOPs/byte
1 33,554,432 33,570,816 1.00
8 268,435,456 33,685,504 7.97
64 2,147,483,648 34,603,008 62.06
512 17,179,869,184 41,943,040 409.60
4096 137,438,953,472 100,663,296 1365.33At batch 1 the ratio is 1.00 FLOP per byte. Every number fetched is used once. Bytes moved barely change from batch 1 to batch 64, because the weight matrix dominates and it is the same matrix either way.
The FLOPs column grows 4096x while the bytes column grows 3x. That gap is the whole opportunity. Batching does not reduce work; it amortises the traffic.
Compare that ratio against your hardware's own ratio of peak arithmetic to peak bandwidth. An NVIDIA RTX A6000 quotes 768 GB/s of memory bandwidth. If a kernel needs fewer FLOPs per byte than the chip can sustain, it will be memory-bound no matter how it is written.
Step two: watch it happen
import torch
assert torch.cuda.is_available()
print("gpu:", torch.cuda.get_device_name(0), "| torch:", torch.__version__)
D = 4096
W = torch.randn(D, D, device="cuda", dtype=torch.bfloat16) # one "layer" of weights
def timed(B, iters=50):
x = torch.randn(B, D, device="cuda", dtype=torch.bfloat16)
for _ in range(10): # warm up: first calls include setup
x @ W
torch.cuda.synchronize()
start, end = torch.cuda.Event(True), torch.cuda.Event(True)
start.record()
for _ in range(iters):
x @ W
end.record()
torch.cuda.synchronize()
ms = start.elapsed_time(end) / iters
flops = 2 * B * D * D
bytes_moved = (B * D + D * D + B * D) * 2
return ms, flops / (ms * 1e-3) / 1e12, bytes_moved / (ms * 1e-3) / 1e9
print(f"{'batch':>6} {'ms/call':>9} {'TFLOP/s':>9} {'GB/s':>8} {'us/token':>9}")
for B in (1, 8, 64, 512, 4096):
ms, tf, gb = timed(B)
print(f"{B:>6} {ms:>9.3f} {tf:>9.1f} {gb:>8.0f} {ms*1000/B:>9.2f}")gpu: NVIDIA RTX A6000 | torch: 2.5.1+cu121
batch ms/call TFLOP/s GB/s us/token
1 0.081 0.4 412 81.39
8 0.064 4.2 528 7.98
64 0.096 22.4 360 1.50
512 0.182 94.1 230 0.36
4096 1.392 98.7 72 0.34These are timings, so they move. Running the same script again on the same machine gave 0.076 ms at batch 1 and 33.5 TFLOP/s at batch 64. Expect swings of tens of percent at small batch sizes, where kernel launch overhead is a large share of a 60-microsecond call. The shape of the table is what reproduces, not the digits.
Read the two middle columns against each other. At batch 1 the card delivers 0.4 TFLOP/s but 412 GB/s. That is over half the quoted 768 GB/s of bandwidth and a rounding error of the available arithmetic. The chip is not slow, it is waiting.
At batch 512 the picture inverts. Arithmetic reaches 94 TFLOP/s and effective bandwidth drops to 230 GB/s. Nothing got faster; the bottleneck moved.
The last column is the business result. Time per token falls from 81 microseconds to 0.36 microseconds, a factor of roughly 200, for the same weights on the same card. This single table explains why serving frameworks work so hard at batching.
Where attention sits in this picture
Two phases, two regimes, and confusing them is the most common performance mistake in LLM work.
| Phase | What runs | Regime |
|---|---|---|
| Prefill (reading your prompt) | Big matrix-matrix products over all prompt tokens | Compute-bound |
| Decode (writing each new token) | Matrix-vector products, plus reading the whole KV cache | Memory-bound |
That split explains the rest of this section. FlashAttention attacks memory traffic inside the attention kernel. Multi-query attention and multi-head latent attention attack the size of the KV cache, which is the dominant read during decode. None of them reduce FLOPs, and none of them need to.
Common mistakes
Optimising FLOPs in a memory-bound kernel. Halving the arithmetic in a kernel that is waiting on memory changes nothing. Measure the achieved bandwidth first.
Benchmarking without warm-up or torch.cuda.synchronize(). CUDA calls are asynchronous. Without a synchronise you are timing how fast Python can queue work, which is a meaningless and impressively fast number.
Comparing TFLOP/s to a marketing peak. Datasheet tensor figures often assume structural sparsity and the friendliest dtype. Compare against a large dense GEMM you measured yourself on the same card.
Assuming a bigger GPU fixes decode latency. Single-stream decoding is bound by memory bandwidth and model size. A card with twice the arithmetic and the same bandwidth gives you close to nothing.
Forgetting the KV cache in the byte count. At long context the cache read per token can exceed the weight read. Then the fix is a smaller cache, not a bigger batch.
Try it yourself
Change D from 4096 to 1024 and rerun. The crossover to compute-bound moves to a larger batch, because a smaller weight matrix carries less traffic to amortise. Then add a torch.float32 run and watch bandwidth halve in effective terms.
What to learn next
- Multi-query attention — the first serious attack on decode-time memory traffic.
- Latency and throughput — the two numbers this trade-off is measured in.
- Is the GPU waiting for data? — the same question, one level up the stack.
Researcher — Mathematics and papers.
The roofline model
From Williams, Waterman and Patterson (2009), Roofline: An Insightful Visual Performance Model for Multicore Architectures, CACM 52(4). Attainable performance for a kernel with arithmetic intensity $I$ (FLOPs per byte of DRAM traffic):
$$ P(I) = \min\left(P_{\max},\; \beta \cdot I\right) $$
$P_{\max}$ is peak arithmetic throughput in FLOP/s and $\beta$ is peak memory bandwidth in bytes/s. The ridge point is
$$ I^{*} = \frac{P_{\max}}{\beta} $$
Kernels with $I < I^{}$ are memory-bound; kernels with $I > I^{}$ are compute-bound. Ridge points have climbed steeply, because arithmetic throughput has grown faster than bandwidth for two decades. An A100 sits near 150 FLOP/byte in bf16; an H100 SXM pairs roughly 989 dense bf16 TFLOP/s with 3.35 TB/s, close to 295 FLOP/byte. Every generation makes it harder to be compute-bound.
Arithmetic intensity of a decoder step
For a dense transformer with hidden size $d$, $L$ layers, $P$ parameters, batch $B$ and current context length $S$, one decoding step reads:
$$ \text{bytes} \approx \underbrace{2P}{\text{weights, bf16}} + \underbrace{2 \cdot 2 \cdot B \cdot S \cdot L \cdot n{kv} \cdot d_h}_{\text{KV cache, bf16}} $$
$n_{kv}$ is the number of key/value heads and $d_h$ the head dimension. The inner factor 2 counts keys and values; the outer 2 is bytes per bf16 element. FLOPs are approximately $2PB$. Hence
$$ I \approx \frac{2PB}{2P + 4BSLn_{kv}d_h} $$
Two limits fall out. With small $B$ and short $S$, $I \to B$: intensity is the batch size, and you are far below any modern ridge point. With large $B$ and long $S$, the KV term dominates and $I \to P / (2SLn_{kv}d_h)$, independent of batch. Batching stops helping once the cache read dominates the weight read, which is the quantitative case for grouped-query attention and cache compression.
Prefill against decode
Prefill processes $S$ tokens at once: FLOPs $\approx 2PS$ against roughly the same weight traffic, so $I \approx S$. At $S = 2048$ this is comfortably compute-bound on any current accelerator.
Decode processes one token per sequence, so $I \approx B$. The gap between the two phases is a factor of $S/B$, and it is the reason disaggregated serving separates prefill and decode onto different hardware pools with different batching policies.
Attention itself
The attention kernel has intensity governed by head dimension, not sequence length. Streaming $Q$, $K$ and $V$ tiles once gives roughly $\Theta(d)$ FLOPs per byte, which for $d = 64$ or $128$ sits below the ridge point of every current data-centre GPU. This is why FlashAttention is a memory-traffic optimisation and why its recomputation in the backward pass is close to free: extra FLOPs cost nothing when the units are idle.
During decode the situation is worse. The query is a single vector, so attention degenerates into a batched matrix-vector product against the entire KV cache. Intensity is order 1, and time is set purely by cache size divided by bandwidth. Grouped-query and latent attention are bandwidth optimisations wearing an architecture costume.
Measurement practice
- Use CUDA events, not
time.time(), and synchronise. Wall-clock timing of asynchronous launches measures the launch queue. - Discard warm-up iterations. The first call includes autotuning, allocator growth and kernel load.
- Report achieved bandwidth alongside achieved FLOP/s. One number alone cannot tell you which roof you are under.
- Prefer Nsight Compute counters (
dram__bytes.sum) over analytic byte counts once caches complicate the picture. Analytic counts assume nothing is reused, which is exactly what an L2 cache exists to falsify.
The trend that matters
Bandwidth per FLOP has fallen for twenty years and continues to fall. The consequence is structural: architectures are increasingly selected for how few bytes they move per token, not how few multiplications they perform. Grouped-query attention, latent attention, sparse attention, mixture-of-experts routing and quantisation are all, at the hardware level, the same idea. Move fewer bytes.
What to learn next
- Multi-query attention — the first serious attack on decode-time memory traffic.
- Latency and throughput — the two numbers this trade-off is measured in.
- Is the GPU waiting for data? — the same question, one level up the stack.