Kontextlängd och kvadratisk kostnad
Kunna beräkna minnes- och tidskostnad för attention som funktion av sekvenslängd.
Förkunskaper
- DAttentionkrävs
- DTidskomplexitet och ordo-notationkrävs
Intuition
Attention jämför varje token med varje annan. Med T tokens blir det T² par — fördubblad kontext ger fyrdubblad kostnad i attention-steget.
| Kontext | Par | Relativt |
|---|---|---|
| 2 k | 4 M | 1× |
| 8 k | 67 M | 16× |
| 32 k | 1,0 G | 256× |
| 128 k | 16 G | 4 096× |
Men två kostnader måste hållas isär:
- Beräkning i prefill: O(T²·d) — växer kvadratiskt.
- Minne under generering: KV-cachen växer linjärt i T, men blir ändå flaskhalsen i praktiken eftersom den ligger kvar hela sessionen.
FlashAttention tog bort den kvadratiska minneskostnaden för attention-matrisen, men inte den kvadratiska beräkningen.
Formellt
FLOPs per lager (ungefärligt, framåtpass):
- Projektioner Q, K, V, O: — linjärt i T.
- Attention-poäng och : — kvadratiskt.
- MLP (4× expansion): — linjärt.
Attention dominerar först när . För betyder det att kontexter under ~4 k tokens domineras av de linjära termerna — vilket förvånar många. Kvadratiskheten biter på allvar först vid långa kontexter.
Minne:
- KV-cache: per sekvens — linjärt i T, men multipliceras med batchen.
- Attention-matrisen: per huvud om den materialiseras. FlashAttention räknar den i block i SRAM och materialiserar den aldrig → minne.
Tre vägar till lång kontext: gles/lokal attention (ändrar vad som beräknas), FlashAttention (ändrar hur det beräknas), och RAG (ändrar om det behöver beräknas alls). Den sista är oftast billigast — hämta de 5 000 relevanta tokens i stället för att mata in 200 000.
Kod
def kostnad(T, d=4096, lager=32, h_kv=8, d_head=128, bytes_per_tal=2):
attn_flops = 2 * T**2 * d * lager
lin_flops = (4 * d**2 + 8 * d**2) * T * lager
kv_gb = 2 * lager * h_kv * d_head * T * bytes_per_tal / 1e9
return {"T": T, "attn_TFLOPs": round(attn_flops / 1e12, 1),
"linjar_TFLOPs": round(lin_flops / 1e12, 1),
"attn_andel": round(attn_flops / (attn_flops + lin_flops), 2),
"kv_cache_GB": round(kv_gb, 2)}
for T in (2_048, 8_192, 32_768, 131_072):
print(kostnad(T))
# {'T': 2048, 'attn_TFLOPs': 1.1, 'linjar_TFLOPs': 13.2, 'attn_andel': 0.08, 'kv_cache_GB': 0.13}
# {'T': 8192, 'attn_TFLOPs': 17.6, 'linjar_TFLOPs': 52.8, 'attn_andel': 0.25, 'kv_cache_GB': 0.54}
# {'T': 32768, 'attn_TFLOPs': 281.5, 'linjar_TFLOPs': 211.1,'attn_andel': 0.57, 'kv_cache_GB': 2.15}
# {'T': 131072, 'attn_TFLOPs': 4504, 'linjar_TFLOPs': 844, 'attn_andel': 0.84, 'kv_cache_GB': 8.59}
Tabellen visar tippningspunkten: vid 2 k är attention 8 % av arbetet, vid 128 k är det 84 %.
Behärskning innebär
- Beräknar attention-kostnaden som funktion av sekvenslängd
- Skiljer beräkning från minne
- Motiverar varför lång kontext är dyrt
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)
- arXiv — FlashAttention — arXiv (öppen åtkomst; licens per artikel)