skip to content

Attention and KV Compression

Few models still run plain multi-head attention: GQA shares key/value heads, MLA compresses them, sparse attention skips them. These choices set the memory math behind long-context serving.

on this pageshow

questions

5

Does FlashAttention change attention's output, and where does its speedup come from?

level: middleimportance: must knowfreq 58%

answer

  1. Same answer, different data movement
  2. Bound by bytes moved, not maths done
  3. Tiles that stay in fast on-chip memory
  4. Running max and running sum
  5. Recompute in backward instead of storing

basics

~20 s

No — FlashAttention computes exactly the same attention, bit-for-bit equivalent up to floating-point reordering. It is faster because it tiles the computation in fast on-chip memory and never writes the full score matrix out to the GPU's main memory, so it moves far less data.

solid answer

~50 s

FlashAttention is an **IO-aware kernel**, not an approximation. A naive implementation computes the full score matrix, writes it to high-bandwidth memory, reads it back to apply softmax, writes again, then reads it once more to weight the values — several round trips over an object whose size grows with the square of the sequence length. FlashAttention instead processes keys and values in tiles that fit in on-chip SRAM, maintaining a running maximum and running sum so the softmax can be computed incrementally and partial outputs rescaled as each tile arrives. The full score matrix therefore never exists in main memory. In the backward pass it recomputes tiles rather than storing them, trading a little arithmetic for a large reduction in memory traffic. The result is the same mathematically, which is why it can be dropped into an existing model with no retraining. Confusing it with sparse or linear attention is the classic error: those change what is computed, FlashAttention changes only where the bytes go.

code

python · 18 lines
python
import numpy as np

scores = np.array([1.0, 3.0, 2.0, 0.5])
values = np.array([10.0, 20.0, 30.0, 40.0])

def tile(s, v):
    m = s.max()
    p = np.exp(s - m)
    return m, p.sum(), (p @ v) / p.sum()

m1, l1, o1 = tile(scores[:2], values[:2])
m2, l2, o2 = tile(scores[2:], values[2:])
m = max(m1, m2)
w1, w2 = l1 * np.exp(m1 - m), l2 * np.exp(m2 - m)
streamed = (w1 * o1 + w2 * o2) / (w1 + w2)

p = np.exp(scores - scores.max())
print(np.allclose(streamed, (p @ values) / p.sum()))

go deeper

for a junior

Know the headline: FlashAttention gives exactly the same result as ordinary attention and is faster because of how it uses GPU memory, not because it skips any work. It needs no retraining.

for a middle

Be ready to explain that attention is memory-bound, that the kernel tiles keys and values through fast on-chip memory, and that an incremental softmax with a running maximum and sum lets partial results be combined exactly.

for a senior

Expect to discuss backward-pass recomputation as a deliberate compute-for-memory trade, and to state clearly that operation count is unchanged so this is not a route to sub-quadratic attention. Knowing it composes with head sharing and sparsity is the practical part.

for a principal

Own the distinction between exact kernel optimisation and approximate architectural change, because teams routinely conflate them when planning long-context work. An efficient kernel is table stakes and buys a constant; changing the mechanism is what changes the scaling, and only the latter is a training-time bet.

## Why a kernel rewrite mattered at all Modern accelerators are enormously faster at arithmetic than at moving data. A GPU can perform hundreds of floating-point operations in the time it takes to fetch one value from its main memory. Attention, implemented naively, is unusually bad on this axis: it produces a large intermediate — the score matrix relating every query to every key — and then reads and writes that intermediate several times. The operation is therefore **memory-bound**, not compute-bound, and speeding up the arithmetic does nothing. FlashAttention's contribution was to notice this and restructure the computation so that the intermediate never leaves fast on-chip memory. ## The memory hierarchy in one paragraph A GPU has a large pool of high-bandwidth memory (HBM) — gigabytes, relatively slow — and a small pool of on-chip SRAM per compute unit — kilobytes, roughly an order of magnitude faster. Anything that must round-trip through HBM costs far more than the same work done inside SRAM. A naive attention kernel writes the scores to HBM, reads them for softmax, writes the normalised probabilities, reads them again to multiply by values. Each of those passes touches an object that grows quadratically with sequence length. ## Tiling and online softmax FlashAttention loads a block of queries and then streams through blocks of keys and values. For each tile it computes that tile's scores directly in SRAM, applies the exponentials there, and accumulates a partial weighted sum of values. The difficulty is that softmax normalises over the *whole* row, and you do not know the row's maximum or its denominator until you have seen every key. The fix is **online softmax**: carry a running maximum and a running sum of exponentials alongside the partial output. When a new tile arrives with a larger maximum, rescale the accumulated output and denominator by the exponential of the difference, then add the new tile's contribution. At the end the accumulator holds exactly the value the one-shot computation would have produced. This is numerically stable, because the running maximum is always subtracted before exponentiating, and it is exact up to the usual reordering effects of floating-point addition. Only the final output — and small per-row statistics — is written to HBM. The quadratic-sized intermediate is never materialised there at all. ## Recomputation in the backward pass Training needs the attention probabilities again to compute gradients. Storing them would reintroduce the very object the forward pass avoided. FlashAttention instead **recomputes** each tile's scores during the backward pass from the saved queries, keys, values and the small saved statistics. That costs extra arithmetic, which is cheap, in exchange for avoiding memory traffic, which is expensive. This is the same trade as gradient checkpointing, applied inside a single operation. ## What it does and does not change It changes: wall-clock time, achievable sequence length within a given memory budget, and the activation memory needed for training. Longer contexts became practical partly because attention stopped requiring a quadratic-sized buffer. It does not change: the model's weights, its outputs (up to floating-point reassociation), its parameter count, or the asymptotic number of arithmetic operations — every query still scores against every key. This is the point people most often get wrong. FlashAttention is not a way to make attention sub-quadratic. If you need to stop doing the work at all, you need a different mechanism: learned sparsity to skip most pairs, or a recurrent formulation that avoids the pairwise structure entirely. FlashAttention makes the exact computation as cheap as it can be; it does not make it a different computation. Later versions refined the same idea for newer hardware — better work partitioning across compute units, fewer non-matrix operations on the critical path, and support for narrower numeric formats — but the principle is unchanged. ## The composability point Because it is exact and weight-preserving, a FlashAttention-class kernel composes with essentially everything else in this area. Head sharing changes how many key/value heads exist; the kernel still tiles over them. Sparse selection changes which blocks are attended; the kernel runs on the selected blocks. As of mid-2026 an efficient tiled kernel is simply assumed in any serious training or serving stack rather than being a named optimisation you choose. ## How to answer Say "exact, not approximate" in the first sentence — that is the discriminator. Then explain memory-bound versus compute-bound, then tiling with online softmax, then backward recomputation. Close by stating explicitly that the operation count is unchanged, which pre-empts the follow-up about whether it makes attention linear.

  • How can softmax be computed correctly if the kernel has only seen part of the row?
    By carrying a running maximum and a running sum of exponentials with the partial output. When a tile arrives whose maximum exceeds the running one, the accumulated output and denominator are rescaled by the exponential of the difference before the new contribution is added. The final accumulator equals the one-shot result, and subtracting the running maximum keeps it numerically stable.
  • Does FlashAttention make attention sub-quadratic in sequence length?
    No. Every query still scores against every key, so the arithmetic still grows quadratically. What becomes near-linear is the memory traffic, because the quadratic intermediate is never written out. To reduce the actual work you need a different mechanism — skipping most query-key pairs through learned sparsity, or replacing the pairwise structure with a recurrent fixed-size state.
  • Why does the backward pass recompute the attention probabilities instead of saving them?
    Saving them would reintroduce the quadratic-sized buffer in main memory that the forward pass was designed to avoid, which is the dominant cost. Recomputing each tile from the saved queries, keys, values and small per-row statistics costs extra arithmetic, which is abundant, to save memory traffic, which is scarce. It is the gradient-checkpointing trade applied inside a single operation.
  • If it is exact, why do outputs sometimes differ slightly from a naive implementation?
    Floating-point addition is not associative, so summing contributions in a different order gives slightly different last bits. Tiling changes the summation order, which produces differences at the level of numerical noise. This is not approximation in the algorithmic sense — no term is dropped or estimated — and it is the same class of discrepancy you see between any two valid reduction orders.

It is the difference between doing a long addition on scratch paper you keep on your desk versus filing each intermediate line in a cabinet down the hall and fetching it back. The sum is identical; the walking is what took all the time.

saying these in an interview costs you the question

  • Says FlashAttention approximates softmax to go faster
  • Claims it makes attention linear in sequence length
  • Thinks the speedup comes from doing fewer multiplications
  • Believes it requires retraining or changes the weights
  • Confuses it with sparse or linear attention mechanisms

context

open as a page

In transformer LLMs, how do GQA and MQA differ from multi-head attention?

level: middleimportance: must knowfreq 72%

basics

~20 s

Multi-head attention gives every query head its own key and value projections. Multi-query attention makes all query heads share a single key/value head. Grouped-query attention sits between them: query heads are split into groups, and each group shares one key/value head.

open as a page

What does Multi-head Latent Attention compress, and how is that different from quantizing the cache?

level: seniorimportance: should knowfreq 42%

basics

~20 s

Multi-head Latent Attention projects each token's keys and values down into one small learned latent vector and retains only that. Per-head keys and values are reconstructed on the fly by up-projection. The compression lives in the trained weights, not in the number format.

open as a page

Why does a trainable sparse-attention selector beat a fixed strided pattern at long context?

level: seniorimportance: should knowfreq 33%

basics

~20 s

A fixed pattern decides which earlier positions each query may see using position alone, so it discards relevant tokens that fall outside the pattern. A trainable selector scores earlier blocks by content and picks per query, and because it is trained jointly the model adapts to the sparsity.

open as a page

When would you interleave linear-recurrent layers with full attention in a long-context model?

level: principalimportance: should knowfreq 26%

basics

~20 s

When the workload needs a very long input but rarely needs exact recall of an arbitrary earlier token. Recurrent layers keep a fixed-size state, so cost per token stops growing with length; the periodic full-attention layers are kept precisely to preserve the exact lookup that a fixed state loses.

open as a page