Part 2 · 1 chapters · ~8 min

Attention and Transformers

Scaled dot-product attention with queries, keys and values, the causal mask, multi-head attention, grouped-query attention, the transformer block (attention, feed-forward, residual connections, normalisation), stacking layers, why attention is quadratic in context length, and mixture-of-experts layers.

3

Attention in a few lines

code
import torch, math
def attention(x, Wq, Wk, Wv):                  # x: (T, d) tokens of one sequence
    Q, K, V = x @ Wq, x @ Wk, x @ Wv
    scores = Q @ K.T / math.sqrt(Q.shape[-1])  # (T, T): every token against every token
    mask = torch.triu(torch.ones_like(scores), diagonal=1).bool()
    scores = scores.masked_fill(mask, float("-inf"))     # causal: no attention to the future
    return torch.softmax(scores, dim=-1) @ V   # weighted sum of values

# one transformer block (pre-norm, as in Llama)
h = x + attention(norm1(x))                    # residual connection
out = h + feed_forward(norm2(h))               # an MLP applied to each token independently
ideawhy
T × T score matrixcompute grows quadratically with context; FlashAttention computes it in tiles without storing it
grouped-query attention (GQA)many query heads share fewer key/value heads: Llama 3.1 8B has 32 query heads and 8 KV heads, shrinking the KV cache 4×
residual connections and normalisationlet gradients flow through dozens of layers (ML maths P5)
mixture of expertsroute each token to a few of many feed-forward experts: more parameters, similar compute per token
ONE ATTENTION HEAD
each token gathers information from earlier tokens
token vectorsQ = xWq, K = xWk, V = xWvscores = Q Kᵀ / √dcausal maskno looking aheadsoftmax per rowattention weightsoutput = weights × V
swipe the figure sideways, or tap expand for full screen
1/4
queries, keys, values
Each token is projected three ways: a query (what am I looking for?), a key (what do I contain?) and a value (what do I pass on?).
three projectionslearned matrices