How Large Language Models Work

Attention Mechanisms Recap


Read this sentence and answer one question: "The trophy would not fit in the brown suitcase because it was too small." What does it refer to?

The suitcase. You knew instantly. Now change one word: "...because it was too large." Now it is the trophy. The pronoun did not move. Nothing about its neighbours changed. The referent flipped because of an adjective eleven tokens away.

Any mechanism that builds a representation of it from a fixed window of nearby words gets this wrong. Any mechanism that compresses the sentence left-to-right into one running summary — the recurrent approach — has to have guessed, before reaching small, that the size of the suitcase would matter. Both designs fail for the same underlying reason: they decide what information to keep before knowing what will be needed.

Attention inverts that. Each position first asks a question, then goes and fetches whatever in the sequence answers it. Nothing is summarised in advance. That is the whole idea, and everything below is mechanism.

"it" resolves differently when one adjective changes0.440.050.070.280.160.090.040.110.580.18trophyfitbrownsuitcaseitit — too smallit — too largeOne row of the score matrix, after dividing by the square root of d_k and taking a softmax.
Nothing about "it" changed; the query it forms is built from a context that did, so its keys match a different word.

Query, key, value — and the arithmetic in full

Every token produces three vectors from its current representation x\mathbf{x}, using three learned weight matrices:

  • Query q=xWQ\mathbf{q} = \mathbf{x}W_Q — "what am I looking for?"
  • Key k=xWK\mathbf{k} = \mathbf{x}W_K — "what do I advertise about myself?"
  • Value v=xWV\mathbf{v} = \mathbf{x}W_V — "what do I actually contribute if selected?"

The separation of key from value matters more than it first appears. A token advertises on one basis and delivers on another. The token cat might advertise "I am the subject noun of this clause" while delivering semantic content about cats. Merging the two would force a single vector to do both jobs.

Let us run the whole computation on three tokens with dk=2d_k = 2, so every number is checkable by hand. Sentence: the cat sat. We compute the output at position 3.

Text
Keys                     Valuesk(the) = [0.2, 0.1]      v(the) = [0.0, 0.1]k(cat) = [0.9, 0.3]      v(cat) = [0.8, 0.6]k(sat) = [0.1, 0.8]      v(sat) = [0.2, 0.9]Query from position 3:   q(sat) = [1.2, 0.4]

Step 1 — raw scores. Dot the query against every key:

Text
q . k(the) = (1.2)(0.2) + (0.4)(0.1) = 0.24 + 0.04 = 0.28q . k(cat) = (1.2)(0.9) + (0.4)(0.3) = 1.08 + 0.12 = 1.20q . k(sat) = (1.2)(0.1) + (0.4)(0.8) = 0.12 + 0.32 = 0.44

Step 2 — scale by dk\sqrt{d_k}. Here 2=1.414\sqrt{2} = 1.414:

Text
0.28 / 1.414 = 0.1981.20 / 1.414 = 0.8490.44 / 1.414 = 0.311

Step 3 — softmax to turn scores into weights that sum to 1:

Text
exp(0.198) = 1.219exp(0.849) = 2.337exp(0.311) = 1.365                sum = 4.921weights: the = 1.219/4.921 = 0.248         cat = 2.337/4.921 = 0.475         sat = 1.365/4.921 = 0.277

Step 4 — weighted sum of values:

Text
out = 0.248 x [0.0, 0.1] + 0.475 x [0.8, 0.6] + 0.277 x [0.2, 0.9]    = [0.000, 0.025] + [0.380, 0.285] + [0.055, 0.249]    = [0.435, 0.559]

The output at position 3 is 47.5% cat, 27.7% itself, 24.8% the. The whole thing in one line:

Attention(Q,K,V)=softmax ⁣(QK⊤dk)V\text{Attention}(Q,K,V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V

Why the dk\sqrt{d_k} is not optional

People treat the scaling factor as a detail. It is load-bearing. If the components of q\mathbf{q} and k\mathbf{k} are roughly independent with unit variance, their dot product over dkd_k dimensions has variance dkd_k — standard deviation dk\sqrt{d_k}. With a realistic dk=64d_k = 64, scores routinely land in the range ±8 or wider. Softmax over such a range is brutal:

ScoressoftmaxWhat happens
[8, 0, −8][0.9997, 0.00034, 0.0000001]Effectively one-hot. Softmax gradient is proportional to p(1−p)p(1-p), which here is about 3×10−43\times10^{-4} — the layer barely learns.
[1, 0, −1][0.665, 0.245, 0.090]A genuine distribution. Gradients flow.

Dividing by dk\sqrt{d_k} pulls the score variance back to about 1, keeping the softmax in the region where it has slope. Remove the scaling and deep models simply fail to train.

Attention weights are not a mysterious learned artefact. They are a softmax over dot products — a similarity search that the model runs against its own sequence, fresh at every position and every layer.

The causal mask: how a model is stopped from cheating

A language model is trained to predict token t+1t+1 from tokens 1..t1..t. If position 3 could attend to position 5, it would see the answer it is meant to predict. Training loss would collapse to nearly zero and the model would learn nothing useful.

The fix is a mask applied to the raw scores before the softmax: set every entry where the key position is later than the query position to −∞-\infty. Since e−∞=0e^{-\infty} = 0, those positions receive exactly zero weight.

Text
Raw scores (4 tokens)          After causal mask[ 2.1  0.3  1.7  0.9 ]         [ 2.1  -inf  -inf  -inf ][ 0.8  3.2  0.4  1.1 ]         [ 0.8   3.2  -inf  -inf ][ 1.4  0.6  2.8  0.2 ]         [ 1.4   0.6   2.8  -inf ][ 0.5  1.9  1.2  2.4 ]         [ 0.5   1.9   1.2   2.4 ]Row 1 after softmax: [1.00, 0,    0,    0   ]Row 2 after softmax: [0.08, 0.92, 0,    0   ]Row 3 after softmax: [0.18, 0.08, 0.74, 0   ]

Row 1 is worth staring at. The first token can only attend to itself, so its softmax is forced to 1.0 regardless of the score. Its attention layer output is just its own value vector — the first position gets no benefit from attention at all.

This mask is also what makes training efficient. Because every row is computed independently in the same matrix operation, a single forward pass over a 2,048-token sequence produces 2,048 separate next-token predictions simultaneously, each conditioned on exactly the right prefix. In one pass you get one sequence's worth of loss signal from every position.

Multiple heads: one comparison is not enough

A single attention operation computes one similarity function. But a token needs several different relationships resolved at once — syntactic agreement, the antecedent of a pronoun, the topic of the paragraph, the matching bracket. Forcing one softmax to serve all of these produces a blurry average of all of them.

Multi-head attention runs hh independent attention operations in parallel over slices of the representation, then concatenates and projects the results:

MultiHead(X)=Concat(head1,…,headh) WO,headi=Attention(XWQi, XWKi, XWVi)\text{MultiHead}(X) = \text{Concat}(\text{head}_1,\dots,\text{head}_h)\,W_O, \qquad \text{head}_i = \text{Attention}(XW_Q^i,\, XW_K^i,\, XW_V^i)

The dimensions are chosen so that total work stays constant. With dmodel=512d_{\text{model}} = 512 and h=8h = 8, each head operates in dk=512/8=64d_k = 512/8 = 64 dimensions. Eight heads of width 64 cost the same as one head of width 512 — you buy diversity of attention patterns for free, at the price of each head having a narrower space to work in.

What individual heads turn out to do

Nobody assigns roles to heads, but interpretability work has found recurring, nameable behaviours:

Head typePatternPurpose
Previous-token headAttends almost entirely to position i−1i-1Copies local context forward; a building block for many circuits
Induction headGiven ...[A][B]...[A], attends from the second [A] to the token that followed the first oneIn-context pattern completion — a large part of why few-shot prompting works at all
Duplicate-token headAttends to earlier occurrences of the same tokenTracking repeated names, variables, list items
Attention sinkDumps most of its weight on the first token, regardless of contentA "do nothing" option — softmax must sum to 1, so a head with nothing to fetch needs somewhere to park its mass

That last row explains an observation that confuses people reading attention maps for the first time: a huge, seemingly meaningless spike on the very first token. It is not a bug. Softmax cannot output all zeros. When a head has no relevant information to retrieve at a given position, it needs a null target, and the first token — usually a fixed start-of-sequence marker with predictable content — becomes that null. This matters practically: several long-context serving tricks work by always keeping the first few tokens in the cache, because evicting the sink destroys the model's output quality far out of proportion to the tokens' apparent importance.

The cost, and what is done about it

Every query is compared against every key. For a sequence of nn tokens that is n2n^2 comparisons, per head, per layer.

Sequence length nnQuery–key pairs (n2n^2)Relative to 512 tokens
512262,1441×
4,09616,777,21664×
32,7681,073,741,8244,096×
131,07217,179,869,18465,536×

Doubling the context quadruples the attention work. This single fact drives most of the engineering below.

FlashAttention: same maths, different memory pattern

A naive implementation writes the full n×nn \times n score matrix to GPU memory, reads it back for the softmax, writes the result, reads it again for the multiply by VV. At n=8,192n = 8{,}192 that intermediate matrix is 67 million entries per head per layer. The bottleneck is not arithmetic — it is moving data between fast on-chip memory and slower high-bandwidth memory.

FlashAttention never materialises the matrix. It processes queries and keys in tiles that fit in on-chip SRAM, computing a running softmax with an online normalisation trick that lets it rescale earlier partial results as new maxima appear. The output is numerically the same attention, but memory use drops from O(n2)O(n^2) to O(n)O(n) and wall-clock time improves severalfold. It is worth being precise about what this does and does not change: the FLOP count is unchanged; only memory traffic is. FlashAttention is not an approximation and does not alter model quality. Later versions (FlashAttention-2 and -3) improve how the work is split across the GPU and exploit newer hardware, and fused kernels of this kind are now what PyTorch's scaled_dot_product_attention and the main serving engines use by default. You rarely call it yourself; you get it unless something forces the slow path — such as asking for the attention weights, as the code later in this lesson does.

MQA and GQA: shrinking the KV cache

During generation, the keys and values of every previous token are cached so they need not be recomputed for each new token. That cache is often the thing that limits how many users a server can handle at once. Work the numbers for a large model with 80 layers, 64 attention heads and head dimension 128, stored in 16-bit precision:

Text
Per token, per layer: 2 (K and V) x 64 heads x 128 dims = 16,384 valuesAcross 80 layers:     16,384 x 80                      = 1,310,720 valuesIn fp16 (2 bytes):    1,310,720 x 2                    = 2.62 MB per tokenA 4,096-token conversation: 2.62 MB x 4,096 = 10.7 GB - for ONE user.

Multi-Query Attention (MQA) keeps all 64 query heads but gives them a single shared key and value head. Grouped-Query Attention (GQA) is the middle ground: group the query heads and give each group its own K/V head. With 8 KV groups instead of 64 heads:

Text
Per token, per layer: 2 x 8 x 128 = 2,048 valuesAcross 80 layers:     163,840 values -> 0.33 MB per token in fp16A 4,096-token conversation: 1.34 GB  - an 8x reduction
Multi-Head (MHA)Grouped-Query (GQA)Multi-Query (MQA)
KV heads (for 64 query heads)6481
KV cache size1×1/81/64
Quality impactBaselineVery smallNoticeable degradation
Where it is usedOlder and smaller modelsThe current default for large modelsSome latency-critical deployments

GQA has become the standard because it captures nearly all the memory saving for nearly none of the quality loss. Note that the parameter count barely changes — the win is entirely in inference-time memory and bandwidth.

Two other ideas attack the same cache from different sides. Sliding-window attention lets most layers see only the last few thousand tokens, so their cache stops growing; models that use it usually interleave these local layers with some full-attention layers so long-range lookups still work. Multi-head latent attention (introduced with DeepSeek-V2) caches one small compressed vector per token and reconstructs per-head keys and values from it, shrinking the cache further than GQA. The question behind all three is the same: how many bytes per token must be kept for every token already seen?

Reading attention maps without fooling yourself

Python
import torchfrom transformers import AutoModelForCausalLM, AutoTokenizertok = AutoTokenizer.from_pretrained("gpt2")model = AutoModelForCausalLM.from_pretrained("gpt2", attn_implementation="eager")inputs = tok("The trophy did not fit in the suitcase because it was too small",             return_tensors="pt")out = model(**inputs, output_attentions=True)# out.attentions is a tuple of length n_layers# each element: (batch, n_heads, seq, seq)attn = out.attentions[5][0, 3]          # layer 6, head 4tokens = tok.convert_ids_to_tokens(inputs["input_ids"][0])it_pos = tokens.index("Ġit")row = attn[it_pos]                       # what "it" attends tofor score, idx in zip(*row.topk(4)):    print(f"{tokens[idx]:>12s}  {score:.3f}")

Two warnings about interpreting the output. First, high attention weight is not the same as high influence. A head can attend heavily to a token whose value vector contributes almost nothing, or attend weakly to one whose value is large. Weight tells you where the model looked, not what it used. Second, a model has dozens of layers and dozens of heads per layer — thousands of these maps. Any pattern you want to find, you can find somewhere. Treat a single striking map as a hypothesis, and test it by intervening: ablate the head and see whether the behaviour actually changes.

What this means when you build with these models

The quadratic cost is the reason long prompts are expensive in a way that surprises people. Doubling your prompt does not double the prefill cost of attention — it quadruples it. If you are stuffing an entire document into a prompt on every request, the cheapest available win is usually to retrieve the relevant 2,000 tokens instead of sending 40,000.

The KV cache is the reason concurrency, not raw compute, is usually what limits a self-hosted deployment. If you are sizing hardware, calculate cache bytes per token from the model's layer count, KV head count and head dimension, multiply by your target context length and concurrent users, and check that against your available memory before you check anything else.

And the attention-sink behaviour is the reason naive context trimming misbehaves. Dropping the oldest tokens from a long conversation sounds harmless, but if the very first tokens are acting as the model's null-attention target, removing them degrades output far more than their content would suggest. Keep the first few tokens pinned and trim from the middle instead.