Hoppa till innehållet
AI-grafen
E· Universitettransformerarkitektur· ca 60 min· utvecklande· verifierad 2026-09-20

Kontextlängd och kvadratisk kostnad

Kunna beräkna minnes- och tidskostnad för attention som funktion av sekvenslängd.

Förkunskaper

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.

KontextParRelativt
2 k4 M1×
8 k67 M16×
32 k1,0 G256×
128 k16 G4 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: 4⋅Td24\cdot T d^2 — linjärt i T.
  • Attention-poäng QK⊤QK^\top och softmax⋅V\text{softmax}\cdot V: 2T2d2T^2 d — kvadratiskt.
  • MLP (4× expansion): 8Td28Td^2 — linjärt.

Attention dominerar först när T≳dT \gtrsim d. För d=4096d = 4096 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: 2LHkvdheadT2 L H_{kv} d_{head} T per sekvens — linjärt i T, men multipliceras med batchen.
  • Attention-matrisen: O(T2)O(T^2) per huvud om den materialiseras. FlashAttention räknar den i block i SRAM och materialiserar den aldrig → O(T)O(T) 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

Alla källor och licenser