Home › What Happens After You Hit Enter › Multi-head self-attention (in the serving loop)
Pipeline stage · Operate

Multi-head self-attention (in the serving loop)

Every token asks a question, every earlier token offers an answer, and the model mixes the answers by how well they match.

In one line

Attention is one operation with two personalities: a compute-bound matmul over the whole prompt, then a memory-bound scan of the entire cache for every single token.

Why you'd careThe thing you have already noticed

You upgraded to a GPU with roughly double the FLOPs, watched prompts process noticeably faster, and saw tokens per second barely move. Or you compared a batch of one crawling at 30 tokens per second against a batch of sixty-four producing over a thousand on the same card. Both are attention showing its two faces. During prefill it multiplies a matrix of queries by a matrix of keys: dense arithmetic that saturates a GPU easily. During decode it has exactly one query row and must read every cached key and value in the model to use it. That second mode moves gigabytes and computes almost nothing, so it is limited by memory bandwidth — and bandwidth is what your new card did not double.

In and outWhat goes in, what comes out

InResidual-stream activations for the tokens being processed — [S, d_model] during prefill, [1, d_model] per sequence during decode — plus this layer's cached keys and values for every previous position.
ProcessProject to Q, K and V; apply RoPE to Q and K; append the new K and V to the cache; score every query against every visible key, scale by 1/sqrt(d_head), apply the causal mask, softmax; take the weighted sum of values; project back to d_model.
OutAn [S, d_model] tensor added into the residual stream, and a cache grown by S positions. The attention weight matrix itself is not returned and, under a fused kernel, is never assembled at all.

The softmax normalises away absolute magnitude, so only the relative ordering of scores survives — a query that matches nothing still emits a full-weight output blended from whatever was least bad, which is the mechanism behind attention sinks. What is lost is the attention map. Modern kernels compute it in tiles and discard it, so logging attention weights at serving time means deliberately running a slower path.

ConceptThe idea underneath

The operation is softmax(Q K^T / sqrt(d_k)) V. Q is one query vector per token, what this position is looking for. K is one key per token, what each position offers. Q K^T dots every query against every key, giving a score matrix of how well each pair matches. Dividing by sqrt(d_k) stops those dot products growing with head dimension and pushing the softmax into saturation. The softmax turns each row of scores into weights summing to 1, and multiplying by V takes a weighted average of what each position actually carries.

Read it as soft, content-addressed retrieval. There are no addresses and no exact matching; a position retrieves a blend of every earlier position weighted by learned similarity. Heads run this independently on lower-dimensional slices, so one head can track subject-verb agreement while another copies a proper noun from four thousand tokens back. The causal mask sets future scores to negative infinity before the softmax, which is what makes the whole thing autoregressive.

The serving fact that matters is arithmetic intensity, FLOPs performed per byte moved. Prefill builds an S-by-S score matrix, so work grows with S squared while weights are read once: hundreds of operations per byte, comfortably compute-bound, 40 to 60 percent of peak achievable. Decode computes one query row against N cached keys, roughly two FLOPs per byte read. The GPU idles almost all the time and the clock is set by how fast the KV cache streams out of HBM. Grouped-query attention, paged caches, quantized KV and continuous batching all exist to attack that second number.

At a glanceSee it

Multi-head self-attention (in the serving loop) diagram

One attention operation, two regimes: a compute-bound prefill matmul and a bandwidth-bound decode step.

The knobsHyperparameters and nuance

  • num_attention_headsH, typically 32 to 128. At fixed width, more heads means more but narrower retrieval channels. Head dimensions below 64 measurably weaken how precisely a query can discriminate between similar keys.
  • head_dimusually 64 or 128, and increasingly declared independently of d_model / H. It sets both discrimination sharpness and per-token cache cost, so it is a quality knob and a memory knob at once.
  • num_key_value_headshow many heads share a cached key/value store. The single largest lever on decode speed, because it divides the bytes read per token directly.
  • softmax scaledefaults to 1/sqrt(head_dim); a few architectures override it, Gemma via query_pre_attn_scalar. A wrong value raises no error, it simply makes every attention distribution uniformly flatter or sharper than the model was trained for.
  • attn_logit_softcappingGemma 2 passes attention logits through a tanh cap before the softmax. Many fast kernels do not support it, so enabling those kernels either hard-errors or silently disables the cap and changes outputs. Model-specific.
  • attention backendvLLM's --attention-backend flag (the older VLLM_ATTENTION_BACKEND environment variable was deprecated in vLLM 0.13.0), or PyTorch's SDPA dispatch via torch.nn.attention.sdpa_kernel. It picks the kernel. It does not change the mathematics, but it changes numerics and can change speed several-fold.

EffectHow this stage moves the answer

Attention is the part of the network that copies. When it works, exact strings survive across thousands of tokens: a UUID quoted back correctly, a variable name reused consistently, a number lifted out of a table without drift. When head capacity is tight — few heads, small head dimension, aggressive KV sharing — that is what degrades first. The model keeps the shape of the answer and loses precision inside it: near-miss identifiers, a plausible but wrong figure, a citation to the right document and the wrong section. From outside it reads as carelessness rather than incapacity, which is why it survives review. Softmax's normalisation contributes directly: because the weights must sum to 1 whether or not anything genuinely matched, a query with no good match still pulls in a blend of whatever scored highest, and the model states that blend with exactly the confidence of a real retrieval.

EvalsWhat it does to your measurements

Every capability routes through attention, so no benchmark isolates it, but two effects appear in measurement constantly. First, the attention variant a model uses determines its long-context scores at fixed parameter count: RULER and needle tests separate MHA, GQA and latent-attention models far more sharply than MMLU does. Second, attention kernels are where batch-dependent nondeterminism enters. Reduction order inside a fused kernel depends on how many sequences are in the batch and how long they are, so the same prompt at temperature 0 can produce different logits and, at a near-tie, a different token. An eval run at concurrency 32 is therefore not bit-comparable with one at concurrency 1, and a regression between two runs is often just a different batch shape. Pin concurrency, kernel backend and dtype before attributing an exact-match delta to the model.

Failure modesWhen it goes wrong

  • Tokens per second falls steadily as a conversation growsevery decode step reads the whole KV cache, so time per token is linear in context length.
  • Doubling tensor parallelism does not double decode throughputattention is bandwidth-bound, and with few key/value heads those heads get replicated rather than split across GPUs.
  • A sharp latency step at one specific context lengththe kernel switched code paths, usually a block-size or head-dimension heuristic crossing a threshold.
  • Long-context answers reference the wrong span while staying fluentretrieval capacity exhausted, not a language failure. Check the attention variant and window before rewriting the prompt.
  • Numerically different outputs from identical weights on two serversdifferent attention backends accumulating in different orders.

PapersWhere this comes from

  • Attention Is All You NeedVaswani et al., 2017. arXiv:1706.03762. Defined scaled dot-product attention and the multi-head split, including the sqrt(d_k) scaling that keeps the softmax out of saturation.
  • Efficiently Scaling Transformer InferencePope et al., 2022. arXiv:2211.05102. Worked out the arithmetic-intensity model that explains why prefill is compute-bound and decode is memory-bound, and derived partitioning strategies serving stacks still follow.
  • Fast Transformer Decoding: One Write-Head is All You NeedShazeer, 2019. arXiv:1911.02150. Identified KV memory bandwidth rather than FLOPs as the decode bottleneck, the observation the entire modern serving stack is built around.
A living map of modern AI — kept current every morning