Hoppa till innehållet
AI-grafen
E· Universitettransformerarkitektur· ca 60 min· grundläggande — ändras sällan· verifierad 2026-09-20

Multi-head attention i detalj

Kunna implementera multi-head attention med masking och KV-cache, och verifiera mot ett referensbibliotek.

Förkunskaper

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

Attn(Q,K,V)=softmax ⁣(QK⊤/dk)V\text{Attn}(Q,K,V) = \text{softmax}\!\big(QK^\top/\sqrt{d_k}\big)V. Skalningen 1/dk1/\sqrt{d_k} håller logits på varians ~1 när q,kq, k har varians 1 per komponent (summan av dkd_k produkter har varians dkd_k); utan den mättas softmax och gradienterna dör. 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, med WiQ,K,V∈Rd×d/hW_i^{Q,K,V}\in\mathbb R^{d\times d/h}. Komplexitet: O(T2d)O(T^2 d) tid, O(T2)O(T^2) minne för poängmatrisen per huvud — FlashAttention undviker att materialisera den. Med KV-cache blir dekodningssteg tt en (1×t)(1\times t)-attention: O(td)O(td).

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

Alla källor och licenser