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.

On this page 9
  1. The short answer
  2. The restaurant you have eaten in
  3. The two situations
  4. Why language models are usually waiting
  5. The one trick that fixes it
  6. A useful number to remember
  7. Where you have already seen this
  8. Remember this
  9. What to learn next

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 hardware

Getting 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

Developer — Code and libraries.

Setup

bash
pip install torch     # the second script needs a CUDA GPU; the first needs nothing

Written 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.

intensity.py
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}")
Output
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.33

At 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

roofline.py
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}")
Output
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.34

These 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.

PhaseWhat runsRegime
Prefill (reading your prompt)Big matrix-matrix products over all prompt tokensCompute-bound
Decode (writing each new token)Matrix-vector products, plus reading the whole KV cacheMemory-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

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