Multi-head attention i detalj
Kunna implementera multi-head attention med masking och KV-cache, och verifiera mot ett referensbibliotek.
Förkunskaper
- DAttentionkrävs
- DTransformers — arkitekturenkrävs
Intuition
Attention i en mening: varje token frågar (Q) alla tokens (K) hur relevanta de är, och hämtar en viktad summa av deras innehåll (V). Vikterna är softmax(QKᵀ/√d).
Multi-head: i stället för en stor attention körs h små parallellt, var och en med egna Q/K/V-projektioner av lägre dimension (d_model/h). Ett huvud kan lära sig «föregående token», ett annat «subjektet i satsen». Utdata konkateneras och projiceras.
Kausal mask: vid generering får token t bara se tokens ≤ t. Man sätter −∞ i logits för framtiden före softmax.
KV-cache: när token t+1 genereras är K och V för alla tidigare tokens oförändrade — spara dem. Då kostar varje ny token O(t) i stället för O(t²). Minnet för cachen (lager × huvuden × t × d) är ofta det som begränsar kontextlängd i praktiken.
Kod
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)
# kausalitet: ändra sista token → första raden oförändrad
x2 = x.copy(); x2[-1] += 1
print(np.allclose(multi_head(x2, *W, h=h)[0], y[0])) # True
Formellt
. Skalningen håller logits på varians ~1 när har varians 1 per komponent (summan av produkter har varians ); utan den mättas softmax och gradienterna dör. Multi-head: , , med . Komplexitet: tid, minne för poängmatrisen per huvud — FlashAttention undviker att materialisera den. Med KV-cache blir dekodningssteg en -attention: .
Behärskning innebär
- Implementerar skalad dot-product attention med kausal mask
- Delar upp i huvuden och slår ihop igen korrekt
- Förklarar KV-cache och varför den gör generering linjär per token
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Attention Is All You Need — arXiv (öppen åtkomst; licens per artikel)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0
- arXiv — FlashAttention — arXiv (öppen åtkomst; licens per artikel)