Skip to content
AI-grafen
EUniversityTransformer architecture· about 60 min· fundamentals that rarely change· verified 2026-09-20· EN

Multi-head attention in detail

Be able to implement multi-head attention with masking and a KV cache, and verify it against a reference library.

Prerequisites

Intuition

Attention in one sentence: every token asks (Q) all the tokens (K) how relevant they are, and fetches a weighted sum of their content (V). The weights are softmax(QKᵀ/√d).

Multi-head: instead of one large attention, h small ones are run in parallel, each with its own Q/K/V projections of a lower dimension (d_model/h). One head can learn «the previous token», another «the subject of the sentence». The outputs are concatenated and projected.

The causal mask: during generation token t may only see tokens ≤ t. You set −∞ in the logits for the future before the softmax.

The KV cache: when token t+1 is generated, K and V for all the earlier tokens are unchanged — save them. Then every new token costs O(t) instead of O(t²). The memory for the cache (layers × heads × t × d) is often what limits the context length in practice.

Code

import numpy as np

def softmax(x, axis=-1):
    x = x - x.max(axis=axis, keepdims=True); e = np.exp(x); return e / e.sum(axis=axis, keepdims=True)

def attention(Q, K, V, causal=False):
    d = Q.shape[-1]
    S = Q @ K.swapaxes(-1, -2) / np.sqrt(d)            # (..., T, T)
    if causal:
        T = S.shape[-1]
        S = np.where(np.tril(np.ones((T, T), bool)), S, -1e9)
    return softmax(S) @ V

def multi_head(x, Wq, Wk, Wv, Wo, h, causal=True):
    T, d = x.shape; dh = d // h
    Q, K, V = x @ Wq, x @ Wk, x @ Wv                   # (T, d)
    split = lambda M: M.reshape(T, h, dh).transpose(1, 0, 2)   # (h, T, dh)
    out = attention(split(Q), split(K), split(V), causal)      # (h, T, dh)
    return out.transpose(1, 0, 2).reshape(T, d) @ Wo

rng = np.random.default_rng(0); d, h, T = 8, 2, 5
W = [rng.normal(0, 0.3, (d, d)) for _ in range(4)]
x = rng.normal(size=(T, d))
y = multi_head(x, *W, h=h)
print(y.shape)                                        # (5, 8)
# causality: change the last token → the first row is unchanged
x2 = x.copy(); x2[-1] += 1
print(np.allclose(multi_head(x2, *W, h=h)[0], y[0]))  # True

Formal

Attn(Q,K,V)=softmax ⁣(QK⊤/dk)V\text{Attn}(Q,K,V) = \text{softmax}\!\big(QK^\top/\sqrt{d_k}\big)V. The scaling 1/dk1/\sqrt{d_k} keeps the logits at a variance of about 1 when q,kq, k have a variance of 1 per component (the sum of dkd_k products has variance dkd_k); without it the softmax saturates and the gradients die. Multi-head: headi=Attn(XWiQ,XWiK,XWiV)\text{head}_i = \text{Attn}(XW_i^Q, XW_i^K, XW_i^V), MHA(X)=[head1;… ;headh]WO\text{MHA}(X) = [\text{head}_1;\dots;\text{head}_h]W^O, with WiQ,K,V∈Rd×d/hW_i^{Q,K,V}\in\mathbb R^{d\times d/h}. The complexity: O(T2d)O(T^2 d) in time, O(T2)O(T^2) in memory for the score matrix per head — FlashAttention avoids materialising it. With a KV cache, decoding step tt becomes a (1×t)(1\times t) attention: O(td)O(td).

Mastery means

  • Implements scaled dot-product attention with a causal mask
  • Splits into heads and merges them back correctly
  • Explains the KV cache and why it makes generation linear per token

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences