Nicht jede Beschleunigung verändert die Architektur
GQA verändert die Zahl geteilter Key/Value-Repräsentationen. FlashAttention berechnet eine dichte Attentionoperation mit einem speichereffizienteren Algorithmus. Ein KV-Cache speichert bereits berechnete Zwischenwerte für spätere Generierungsschritte. Diese Techniken lösen verwandte, aber verschiedene Engpässe.
Wenn jemand „faster Attention“ sagt, frage daher zuerst: weniger theoretische Rechenarbeit, weniger gespeicherte Zwischenwerte, weniger Speicherverkehr oder weniger redundante Rechnungen? Ein Benchmark auf einer Hardware beantwortet nicht automatisch alle diese Fragen.
MHA, MQA und GQA
Bei Multi-Head Attention hat jeder Query-Head einen eigenen Key- und Value-Head. Bei Multi-Query Attention teilen alle Query-Heads ein einziges Key/Value-Paar. Grouped-Query Attention liegt dazwischen: Gruppen von Query-Heads teilen jeweils einen Key- und Value-Head.
Wenn und , benutzen jeweils vier Query-Heads dieselben K/V-Repräsentationen. Die Q-Projektion liefert weiterhin Komponenten, K und V nur jeweils . Das reduziert die K/V-Projektionsgewichte und besonders den Cachebedarf. GQA-Originalarbeit.
GQA im leicht lesbaren Code
repeat = cfg.heads // cfg.kv_heads
k_for_attention = k.repeat_interleave(repeat, dim=1)
v_for_attention = v.repeat_interleave(repeat, dim=1)
out = F.scaled_dot_product_attention(q, k_for_attention, v_for_attention, is_causal=True)Diese explizite Wiederholung ist leicht nachvollziehbar, materialisiert aber zusätzliche Tensoren. Eine optimierte GQA-Implementierung kann das vermeiden. Unser Lehrprojekt spart Projektionsparameter, beansprucht aber nicht, durch diese naive Wiederholung bereits den optimalen Kernel-Speicherbedarf zu erreichen.
Prefill und Decode
Beim Prefill liest das Modell den vollständigen Prompt. Viele Query-Positionen werden parallel verarbeitet. Beim Decode wird jeweils ein neuer Token ergänzt. Ohne Cache rechnest du sämtliche K/V-Repräsentationen des bisherigen Kontexts immer wieder aus.
Mit einem Cache speichert jede Schicht die Keys und Values der bisherigen Positionen. Der neue Token erzeugt nur seine neuen Q/K/V-Vektoren. Seine Query liest die gespeicherten Keys und Values plus den neuen Eintrag. Die alten Hidden States müssen nicht neu berechnet werden, solange Kontext und Modellregeln unverändert fortgeführt werden.
Cachebedarf berechnen
Für Layerzahl L, Batch B, gespeicherte Länge T, K/V-Kopfzahl Hkv, Kopfdimension d und Bytes pro Element e ergibt sich ungefähr:
M_KV ist der benötigte Speicher in Bytes, also Speichereinheiten. L zählt die Schichten; B die gleichzeitig verarbeiteten Beispiele; T die gespeicherten Positionen; H_kv die verschiedenen Key/Value-Köpfe; d die Komponenten pro Kopf; e die Bytes pro gespeicherter Zahl. Alle Faktoren werden multipliziert. Die Zwei steht für die zwei getrennten Felder K und V. Bei L=6, B=1, T=1024, Hkv=2, d=64 und zwei Bytes pro Element ergeben sich genau 2 × 6 × 1 × 1024 × 2 × 64 × 2 = 3,145,728 Bytes, also drei MiB. Ein MiB (Mebibyte) sind 1024 × 1024 = 1,048,576 Bytes. Zusätzlich gibt es Gewichte, temporäre Tensoren und Implementierungsmetadaten.
Batchgröße und Kontextlänge wachsen linear in diese Cacheformel ein. Das erklärt, weshalb ein Modell mit kleinen quantisierten Gewichten bei vielen parallelen langen Gesprächen trotzdem erheblichen Arbeitsspeicher benötigen kann.
Der wichtige Maskenfehler beim Cache
Beim normalen vollständigen Forward sind Query- und Key-Längen gleich. Beim Decode mit Cache kann die Query-Länge eins und die Key-Länge tausend sein. Ein blindes is_causal=True mit einer oben links ausgerichteten Dreiecksmaske kann dann die falschen Keys erlauben. Die neue Query soll alle bisherigen Keys sehen, nicht nur den ersten.
Für einen einzelnen neuen Token bei vollständig vergangenem Cache kann es richtig sein, keine zusätzliche kausale Einschränkung anzuwenden. Für mehrere neue Tokens braucht die Maske einen korrekten Positionsoffset. Ebenso muss RoPE die tatsächliche neue Position verwenden. Unser Projekt hat absichtlich noch keinen persistenten KV-Cache; eine Cacheerweiterung ist eine fortgeschrittene Aufgabe mit einem Vergleichstest gegen den ungecachten Decoder.
FlashAttention
Naive Attention speichert die vollständige Score- oder Gewichtsmatrix. FlashAttention verarbeitet Blöcke und verwendet eine numerisch stabile Online-Softmax, um das dichte Ergebnis zu berechnen, ohne die gesamte Matrix im großen GPU-Speicher zu materialisieren. Die dichte Paararbeit bleibt grundsätzlich vorhanden; das Verfahren ist keine automatische lineare Sparse Attention. FlashAttention.
PyTorch SDPA kann abhängig von Gerät, Dtype, Masken und Form einen passenden Kernel auswählen. Ein Aufruf dieser Funktion beweist nicht, dass FlashAttention auf deinem Mac oder jeder CUDA-Konfiguration tatsächlich gewählt wurde. Miss und prüfe das Backend, statt einen Funktionsnamen als Leistungsgarantie zu lesen.
Zwei weitere verbreitete Verfahren: Multi-Head Latent Attention (MLA) speichert statt K und V einen komprimierten latenten Vektor pro Token und verkleinert so den Cache stärker als GQA (DeepSeek-V2). Speculative Decoding lässt ein kleines Entwurfsmodell (oder zusätzliche Vorhersageköpfe) mehrere Tokens vorschlagen, die das große Modell in einem einzigen Forward prüft. Bei korrekter Annahmeregel bleibt die Ausgabeverteilung unverändert (Leviathan et al., 2022).
Dein Erweiterungsplan
Implementiere zuerst einen Cache für den klassischen Decoder ohne Kontextverschiebung. Vergleiche bei jeder Position die ungecachten und gecachten Logits. Erweitere dann RoPE mit einem Offset und wiederhole den Vergleich. Erst danach untersuche GQA-Caches und lange Kontexte. Ein Cache, der schneller ist und andere Tokens vorhersagt, kann schlicht falsch sein.