Attention je jediná část modelu, kde tokeny mluví mezi sebou a zároveň jediná část, jejíž náklady rostou s délkou kontextu. Proto se za posledních pár let většina architektonických vylepšení točí právě kolem ní. Tahle kapitola vysvětlí KV cache, která je hlavním spotřebitelem paměti při provozu, a cestu MHA → MQA → GQA → MLA, kterou prošly všechny moderní modely.
KV cache: proč generování není kvadratické
Naivně by se při každém novém tokenu musela přepočítat pozornost pro celý dosavadní text. To by bylo neúnosné. Trik je, že K a V dřívějších tokenů se nemění, jednou spočítané se dají uložit a znovu použít.
prefill (zpracování promptu): spočítej K,V pro všech N tokenů naráz → ulož do cache
decode (každý další token): spočítej Q,K,V jen pro nový token
načti K,V z cache, spočítej pozornost, přidej nové K,V
Díky tomu je generování jednoho tokenu lineární v délce kontextu, ne kvadratické. Cenou je paměť, a ta je velká.
Vzorec, který by měl AI engineer znát zpaměti
KV cache = 2 × n_layers × n_kv_heads × head_dim × seq_len × batch × bajty
↑
K a V
Příklad, model 32 vrstev, 32 KV hlav, head_dim 128, kontext 8 192, jeden uživatel, FP16:
2 × 32 × 32 × 128 × 8192 × 1 × 2 B ≈ 4,3 GB
Čtyři gigabajty jen na cache jednoho uživatele, k tomu 14 GB na samotný model. A teď si představ padesát souběžných konverzací. Tohle je důvod, proč se serving modelů řeší tak intenzivně a proč vzniklo GQA.
MHA → MQA → GQA
Multi-Head Attention (MHA, 2017). Každá hlava má vlastní Q, K i V. Nejvyšší kvalita, největší cache.
Multi-Query Attention (MQA, 2019). Všechny hlavy sdílejí jediné K a V; liší se jen Q.
Cache se zmenší n_heads-krát (třeba 32×), ale kvalita znatelně klesne.
Grouped-Query Attention (GQA, 2023). Kompromis, který dnes používá skoro každý: hlavy se rozdělí do skupin a každá skupina sdílí jedno K a V.
MHA: 32 hlav Q, 32 hlav K/V → cache 100 %, kvalita 100 %
GQA: 32 hlav Q, 8 hlav K/V → cache 25 %, kvalita ~99 %
MQA: 32 hlav Q, 1 hlava K/V → cache 3 %, kvalita ~95 %
Proto ve specifikaci vidíš n_heads: 32, n_kv_heads: 8. Ten druhý údaj ti přímo říká, kolikrát
menší bude cache a je to při odhadu paměti důležitější číslo než počet hlav.
Multi-head Latent Attention (MLA). Novější přístup, kde se K a V nejdřív zkomprimují do menšího latentního vektoru a při použití se rozbalí. Cache klesne ještě víc než u GQA při zachování kvality blízké MHA. Za cenu složitější implementace.
FlashAttention: stejná matematika, jiná práce s pamětí
Naivní implementace vytvoří v paměti celou matici skóre seq_len × seq_len. Pro 8 k kontext to
je 64 milionů čísel na hlavu a vrstvu a hlavně obrovský přenos mezi rychlou SRAM a pomalou
HBM pamětí grafiky.
FlashAttention počítá totéž po dlaždicích tak, aby se úplná matice nikdy nemusela materializovat. Výsledek je matematicky stejný, ale několikanásobně rychlejší a s dramaticky nižší spotřebou paměti. Dnes je to standard ve všech serverovacích knihovnách a je to hezká ukázka toho, že u velkých modelů rozhoduje pohyb dat, ne počet operací.
Jak se zkrotil dlouhý kontext
Plná pozornost přes 200 000 tokenů je i s FlashAttention drahá, tak se používají triky:
- Sliding window. Token vidí jen posledních N tokenů (třeba 4 096). Delší vazby se přenesou napříč vrstvami, protože každá vrstva okno posune.
- Střídání vrstev. Část vrstev má plnou pozornost, část jen okno. Kompromis mezi kvalitou a cenou.
- Sink tokeny. Ukazuje se, že modely „kotví“ pozornost na prvních pár tokenech. Když se při ořezávání ponechají, kvalita dlouhého kontextu se výrazně zlepší.
- Komprese KV cache. Kvantizace cache na 8 nebo 4 bity, případně zahazování méně důležitých pozic.
Co si z toho odnést do praxe
- Při odhadu paměti počítej model i cache. Cache často přeroste model, když má víc uživatelů.
n_kv_headsje důležitější nežn_heads, když plánuješ hardware.- Dlouhý kontext je drahý dvakrát: jednou při prefillu (kvadraticky), pak trvale v cache.
- Kvantizace KV cache je nejrychlejší způsob, jak zvýšit počet souběžných uživatelů, když ti dochází VRAM.
Cvičení
- Model má 80 vrstev, 8 KV hlav,
head_dim128. Kolik zabere KV cache pro kontext 32 k tokenů ve fp16, pro jednoho uživatele? A pro deset souběžných? - Kdyby tentýž model používal plné MHA se 64 hlavami, kolikrát větší by cache byla?
- Máš kartu s 80 GB, model ve fp16 zabere 40 GB. Kolik souběžných konverzací na 8k kontextu uneseš při GQA z bodu 1?
Náčrt řešení: rozbal, až si cvičení zkusíš sám
- 10 GiB na jednoho uživatele, 100 GiB na deset. Dosazení do
2 × vrstvy × kv_hlavy × head_dim × délka × bajty, tedy2 × 80 × 8 × 128 × 32768 × 2 B = 10 737 418 240 B. Dvojka na začátku je za K a V zvlášť. Deset souběžných uživatelů se do žádné jedné karty nevejde, a to je celý důvod, proč existuje paged attention a proč se s dlouhým kontextem počítá jako s nákladovou položkou. - Osmkrát, tedy 80 GiB. Cache roste lineárně s počtem KV hlav,
64 / 8 = 8. GQA je tedy osminásobná úspora paměti při prakticky nezměněné kvalitě, což je jeden z nejlepších poměrů cena ku výkonu v celé moderní architektuře. - Zhruba šestnáct. Na 8k kontextu je cache
2 × 80 × 8 × 128 × 8192 × 2 B = 2,5 GiB, po odečtení modelu zbývá asi 40 GB, tedy40 / 2,5 ≈ 16. V praxi počítej s méně, protože část paměti spolykají aktivace a fragmentace.
Kam dál
KV cache je hlavní důvod, proč se u modelu počítá paměť jinak než u klasického programu. Pokud tě zajímá, jak se to promítne do provozu, pokračuj na inference stack a čtení specifikace modelu. Mechaniku samotné pozornosti krok za krokem najdeš v attention ručně na číslech.
Shrnutí
- KV cache dělá generování lineárním, ale spotřebuje obrovské množství paměti.
- GQA (sdílené K/V mezi skupinami hlav) je dnešní standard: čtvrtinová cache, téměř žádná ztráta.
- FlashAttention nemění matematiku, jen pohyb dat a tím zrychluje a šetří paměť.
- Dlouhý kontext se řeší okny, střídáním vrstev, sink tokeny a kvantizací cache.
Proč není generování tokenu kvadraticky drahé, když attention porovnává každý s každým?
Protože klíče a hodnoty dřívějších tokenů se nemění a drží se v KV cache. Pro nový token se počítá jen jeho Q, K a V a pozornost proti uložené cache, tedy lineárně v délce kontextu. Kvadratický je jen prefill, tedy prvotní zpracování promptu.
Model má n_heads 32 a n_kv_heads 8. Co to znamená pro paměť?
Že jde o GQA: osm skupin sdílí klíče a hodnoty, takže KV cache je čtvrtinová oproti plnému MHA při prakticky stejné kvalitě. Pro odhad paměti je rozhodující právě n_kv_heads, ne n_heads.
Čím se liší FlashAttention od běžné implementace?
Počítá matematicky totéž, ale po dlaždicích, aniž by kdy vytvořil celou matici skóre v paměti. Šetří tím hlavně přenosy mezi rychlou a pomalou pamětí GPU, což je u velkých modelů skutečné úzké hrdlo, proto je znatelně rychlejší a paměťově úspornější.
