Vzorec
softmax(QKᵀ / √d) · Vvypadá jako zaklínadlo, dokud si ho jednou nespočítáš. Tak si ho spočítáme, na třech tokenech a čtyřech rozměrech, s čísly, která si můžeš ověřit na papíře. Až tohle uděláš, přestane být attention magie a stane se z ní obyčejné násobení matic s jedním chytrým krokem uprostřed.
Zadání
Věta: „kočka spí venku“, tedy tři tokeny. Zvolíme si směšně malý model, aby se čísla vešla na obrazovku:
d_model = 4 (každý token je vektor o čtyřech číslech)
head_dim = 4 (jedna hlava, bere celý vektor)
Po embeddingu máme tři vektory (čísla jsou vymyšlená, ale realistická velikostí):
x₁ (kočka) = [ 1.0, 0.0, 1.0, 0.0 ]
x₂ (spí) = [ 0.0, 1.0, 0.0, 1.0 ]
x₃ (venku) = [ 1.0, 1.0, 0.0, 0.0 ]
Krok 1: tři projekce
Model má tři naučené matice W_Q, W_K, W_V (každá 4 × 4). Pro jednoduchost vezmeme takové,
které jen vybírají a míchají složky:
q₁ = x₁ · W_Q = [ 1.0, 0.0, 0.5, 0.0 ]
k₁ = x₁ · W_K = [ 0.9, 0.1, 0.4, 0.0 ]
v₁ = x₁ · W_V = [ 0.2, 0.8, 0.0, 0.1 ]
k₂ = x₂ · W_K = [ 0.0, 1.0, 0.0, 0.9 ]
v₂ = x₂ · W_V = [ 0.7, 0.1, 0.3, 0.0 ]
k₃ = x₃ · W_K = [ 0.8, 0.9, 0.1, 0.0 ]
v₃ = x₃ · W_V = [ 0.1, 0.2, 0.9, 0.4 ]
Význam těch tří rolí:
- q (query): „co tenhle token hledá“
- k (key): „co tenhle token nabízí ostatním“
- v (value): „co předá tomu, kdo si ho vybere“
Krok 2: skóre = skalární součin
Počítáme pozornost pro třetí token („venku“). Jeho q₃ porovnáme s klíči všech tokenů,
na které smí vidět (tedy 1, 2 a 3, na sebe ano, do budoucnosti ne).
Řekněme q₃ = [ 0.9, 0.8, 0.2, 0.1 ]. Skalární součin je prosté „vynásob po složkách a sečti“:
q₃ · k₁ = 0.9×0.9 + 0.8×0.1 + 0.2×0.4 + 0.1×0.0 = 0.81 + 0.08 + 0.08 + 0 = 0.97
q₃ · k₂ = 0.9×0.0 + 0.8×1.0 + 0.2×0.0 + 0.1×0.9 = 0 + 0.80 + 0 + 0.09 = 0.89
q₃ · k₃ = 0.9×0.8 + 0.8×0.9 + 0.2×0.1 + 0.1×0.0 = 0.72 + 0.72 + 0.02 + 0 = 1.46
Vyšší číslo = větší podobnost mezi tím, co token hledá, a tím, co druhý nabízí.
Krok 3: dělení odmocninou a proč
√head_dim = √4 = 2
0.97 / 2 = 0.485
0.89 / 2 = 0.445
1.46 / 2 = 0.730
Proč se vůbec dělí? Skalární součin sčítá head_dim součinů. Čím delší vektory, tím větší
čísla vycházejí, u head_dim = 128 klidně v desítkách. A pak nastane tohle:
skóre [10, 2, 1] → softmax → [0.9997, 0.0003, 0.0001]
Softmax se saturuje: jeden token dostane skoro všechnu váhu a ostatní prakticky nulu. To má
dva zlé důsledky: model se dívá jen na jedno místo a gradient v saturované oblasti je téměř
nulový, takže se přestane učit. Dělení √d udrží skóre v rozumném rozsahu bez ohledu na to,
jak velké hlavy zvolíš.
Krok 4: softmax
Softmax převede libovolná čísla na pravděpodobnosti, nezáporné a v součtu jedna:
softmax(zᵢ) = e^zᵢ / Σ e^zⱼ
e^0.485 = 1.624
e^0.445 = 1.561
e^0.730 = 2.075
součet = 5.260
váhy = [ 1.624/5.260, 1.561/5.260, 2.075/5.260 ]
= [ 0.309, 0.297, 0.394 ] (součet = 1.000 ✓)
Token „venku“ tedy věnuje 31 % pozornosti slovu „kočka“, 30 % slovu „spí“ a 39 % sám sobě.
💡 Exponenciála tam není náhodou: zvětšuje rozdíly (větší skóre dostane nepoměrně větší váhu) a zaručuje kladné hodnoty. Zároveň je to funkce, která se hezky derivuje, což potřebuješ při backpropu.
Krok 5: kauzální maska
Kdybychom počítali pozornost pro první token, nesmí vidět na druhý a třetí. Řeší se to tak, že se jejich skóre před softmaxem nastaví na minus nekonečno:
skóre pro token 1: [ 0.485, −∞, −∞ ]
e^(−∞) = 0
softmax → [ 1.000, 0.000, 0.000 ]
Elegantní: maska se aplikuje na skóre, softmax se o zbytek postará sám. Proto se v kódu píše
scores.masked_fill(mask, float('-inf')).
Krok 6: vážený součet hodnot
Výstup pozornosti pro token 3 je vážený průměr hodnot podle spočítaných vah:
out₃ = 0.309 × v₁ + 0.297 × v₂ + 0.394 × v₃
= 0.309 × [0.2, 0.8, 0.0, 0.1]
+ 0.297 × [0.7, 0.1, 0.3, 0.0]
+ 0.394 × [0.1, 0.2, 0.9, 0.4]
složka 1: 0.309×0.2 + 0.297×0.7 + 0.394×0.1 = 0.062 + 0.208 + 0.039 = 0.309
složka 2: 0.309×0.8 + 0.297×0.1 + 0.394×0.2 = 0.247 + 0.030 + 0.079 = 0.356
složka 3: 0.309×0.0 + 0.297×0.3 + 0.394×0.9 = 0 + 0.089 + 0.355 = 0.444
složka 4: 0.309×0.1 + 0.297×0.0 + 0.394×0.4 = 0.031 + 0 + 0.158 = 0.189
out₃ = [ 0.309, 0.356, 0.444, 0.189 ]
A to je celá attention. Skalární součiny, dělení odmocninou, softmax, vážený průměr. Nic víc v tom není.
Krok 7: víc hlav a co je W_O
Reálný model tohle nedělá jednou, ale n_heads-krát paralelně nad různými částmi vektoru:
d_model = 4096, n_heads = 32 → head_dim = 128
q, k, v mají 4096 čísel → rozřežou se na 32 kousků po 128
každá hlava počítá kroky 2–6 nezávisle nad svým kouskem
výstupy hlav (32 × 128) se poskládají zpět vedle sebe → 4096
A pak přijde krok, který se v popisech často vynechává: výstupní projekce W_O
(matice 4096 × 4096). Spojené výstupy hlav se jí ještě promítnou, teprve výsledek se přičte
do residual streamu.
out = concat(hlava₁, …, hlava₃₂) · W_O
W_O je důležitá: umožňuje modelu míchat informace mezi hlavami a rozhodnout, kolik z čeho
se do residual streamu zapíše. Je to čtvrtina parametrů attention části.
Rozměry na jednom místě
Pro dávku B sekvencí délky T:
x [B, T, d_model] vstup
q, k, v [B, T, d_model] po projekci
[B, n_heads, T, head_dim] po rozdělení na hlavy
skóre [B, n_heads, T, T] každý s každým ← tady je ta kvadratika
váhy [B, n_heads, T, T] po softmaxu
out [B, n_heads, T, head_dim]
[B, T, d_model] po spojení hlav
[B, T, d_model] po W_O
Matice [B, n_heads, T, T] je ta, kterou FlashAttention nikdy
nevytvoří celou v paměti a teď už vidíš proč: pro T = 8192 a 32 hlav je to 2 miliardy čísel
na jednu vrstvu.
Cvičení
- Spočítej
out₁(pozornost pro první token). Nápověda: po maskování má váhu 1,0 sám na sebe, takže výsledek je přesněv₁. Rozumíš, proč to tak vyjde? - Zvětši všechna skóre desetinásobně (jako by chybělo dělení
√d) a spočítej softmax znovu. O kolik klesne váha nejslabšího tokenu? - Model má
d_model = 4096,n_heads = 32. Kolik čísel má matice skóre pro sekvenci 1 024 tokenů a jednu vrstvu? A kolik gigabajtů to je ve fp16?
Náčrt řešení: rozbal, až si cvičení zkusíš sám
- Vyjde přesně
v₁. První token má po maskování jediné nemaskované skóre, softmax z jednoho čísla je vždy1,0, a vážený průměr s jedinou vahou jedna je ta hodnota samotná. Pointa je hlubší, než se zdá: první token nemá z čeho čerpat kromě sebe, takže attention pro něj nedělá vůbec nic. Užitečnost pozornosti roste s tím, kolik je vlevo kontextu. - Nejslabší váha klesne z
0,297na0,051, tedy zhruba šestkrát. Skóre[4,85, 4,45, 7,30]dají po softmaxu[0,075, 0,051, 0,874]. Nejsilnější token si mezitím polepšil z0,395na0,874. Přesně tohle dělá chybějící dělení odmocninou: softmax se blíží k tomu, že vybere jeden token a na zbytek se vykašle. A protože je v saturované oblasti, gradient je skoro nulový a model se z toho nevyhrabe. - 33 554 432 čísel, ve fp16 tedy 64 MiB na jednu vrstvu. Výpočet je
n_heads × T × T, tedy32 × 1024 × 1024. Všimni si, co v tom čísle není:d_model. Velikost matice skóre nad_modelvůbec nezávisí, roste jen s počtem hlav a s druhou mocninou délky kontextu. Proto se plná matice v praxi nikdy nevytváří naráz, viz flash attention.
Shrnutí
- Attention = skalární součiny (skóre) → dělení
√d→ softmax → vážený průměr hodnot. - Dělení odmocninou brání saturaci softmaxu, která by zastavila učení.
- Kauzální maska se aplikuje před softmaxem nastavením skóre na minus nekonečno.
- Víc hlav počítá totéž nad kousky vektoru;
W_Ovýsledky promíchá a zapíše do residual streamu.
Proč se skóre dělí odmocninou z head_dim?
Skalární součin sčítá head_dim součinů, takže s rostoucí velikostí hlavy rostou i skóre. Velká skóre saturují softmax, jeden token dostane skoro všechnu váhu a gradienty se blíží nule, takže se model přestane učit. Dělení odmocninou drží rozsah skóre stabilní nezávisle na velikosti hlavy.
Jak se technicky zajistí, že token nevidí do budoucnosti?
Skóre k pozdějším tokenům se před softmaxem nastaví na minus nekonečno. Exponenciála z minus nekonečna je nula, takže po normalizaci mají tyhle pozice přesně nulovou váhu, maska se tedy řeší jedním krokem před softmaxem, ne dodatečným filtrováním.
K čemu slouží matice W_O, když jsou výstupy hlav už spočítané?
Spojené výstupy hlav se jí promítnou zpět do rozměru modelu. Umožňuje míchat informace mezi hlavami a rozhodnout, co a v jaké míře se zapíše do residual streamu. Je to zhruba čtvrtina parametrů attention části, takže rozhodně nejde o formalitu.
