Tokenwerk · Das LLM-Lehrbuch

Kapitel 18 · III · Den Transformer bauen · 5 Minuten

Kausalität und mehrere Köpfe

Wie das Modell parallel trainiert, ohne die Antwort zu sehen. Und wie aus einer Attentionoperation mehrere Köpfe werden.

Die nächste Antwort darf noch nicht sichtbar sein

Wir kennen jetzt die Attention-Tabelle: Eine Zeile sagt, wie eine Position verschiedene Quellen mischt. Beim Training liegt der gesamte Beispielsatz schon im Speicher. Trotzdem darf eine Position nur sich selbst und frühere Positionen lesen. Sonst könnte sie das gesuchte nächste Token einfach abschreiben.

Für drei Positionen sind die erlaubten Quellen: Position 1 liest nur 1; Position 2 liest 1 und 2; Position 3 liest 1, 2 und 3. Eine Maske ist eine Tabelle, die genau diese Erlaubnisse festhält. „Kausal“ heißt hier: Eine Berechnung hängt nicht von späteren Textpositionen ab.

Um eine Quelle auszuschließen, setzt man ihren Score vor Softmax auf minus unendlich, geschrieben −∞. Das ist eine mathematische Grenzschreibweise: Die Exponentialfunktion nähert sich für immer negativere Zahlen null. So bekommt die verbotene Quelle Gewicht null. Im Code kann ein Wahrheitswert wie True (wahr) oder False (falsch) die Erlaubnis beschreiben; welche Richtung „erlaubt“ bedeutet, hängt von der konkreten Funktion ab.

Parallel rechnen, ohne in die Zukunft zu sehen

Beim Training sind alle Tokens einer Sequenz als Array vorhanden. Trotzdem darf die Ausgabe an Position tt nur Eingabepositionen bis einschließlich tt lesen. Ihr Target ist xt+1x_{t+1}; das steht eine Position weiter rechts. Die kausale Maske verbietet genau diesen Blick nach rechts.

Das bedeutet nicht, dass wir jedes Präfix in einer eigenen Schleife vorwärts rechnen müssen. Wir berechnen viele Query-Zeilen gleichzeitig und maskieren die unerlaubten Einträge. So erhält das Modell an jeder Position ein korrekt eingeschränktes Präfix und kann alle zugehörigen Next-Token-Ziele parallel lernen.

Warum die Diagonale erlaubt ist

Die Eingabe an Position tt ist xtx_t. Das zugehörige Ziel ist xt+1x_{t+1}. Das aktuelle Eingabetoken gehört also zum erlaubten Kontext. Eine Maske, die auch die Diagonale verbietet, lässt der ersten Position keine einzige erlaubte Quelle: Alle Scores wären −∞, und Softmax ergäbe keine gültige Verteilung (in PyTorch NaN). Unser Standarddecoder erlaubt die Diagonale.

python
T = q.size(-2)
allowed = torch.ones(T, T, dtype=torch.bool, device=q.device).tril()
scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
scores = scores.masked_fill(~allowed, float('-inf'))
weights = scores.softmax(dim=-1)
out = weights @ v

Verbotene Scores werden auf minus unendlich gesetzt. Die Exponentialfunktion liefert dafür null. Ein bloßer Score von null wäre falsch: Softmax kann ihm eine positive Wahrscheinlichkeit zuweisen.

Eine gefährliche Maskenkonvention

Unterschiedliche PyTorch-APIs verwenden boolesche Masken unterschiedlich. Bei scaled_dot_product_attention bedeutet eine wahre boolesche Maskenposition „darf teilnehmen“. Bei anderen Attention-Interfaces kann wahr „blockiert“ bedeuten. Kopiere eine Maske deshalb nicht blind zwischen APIs.

Für den vollständigen gleichlangen Decoder-Forward benutzen wir:

python
out = F.scaled_dot_product_attention(
    q, k, v,
    dropout_p=0.0,
    is_causal=True,
)

Die Dropoutwahrscheinlichkeit musst du ausdrücklich auf null setzen, wenn du keinen Dropout willst. Die funktionale API leitet das nicht automatisch aus model.eval() ab. Das Referenzmodell setzt generell keinen Attention-Dropout ein.

Multi-Head Attention

D ist weiterhin die gesamte Darstellungsbreite. H ist die Anzahl paralleler Attention-Köpfe, und kleines d die Breite eines einzelnen Kopfes. Bei D = 128 und H = 4 ist d = 32. Anstatt einen Kopf mit Breite DD zu benutzen, verwenden wir HH Köpfe mit Breite d=D/Hd=D/H. Jeder Kopf besitzt seine eigene Query-, Key- und Value-Projektion. Praktisch können wir alle Köpfe in einer großen linearen Projektion berechnen und danach die Kopfachse herausformen.

python
q = self.q_proj(x)  # [B,T,D]
q = q.view(B, T, H, d).transpose(1, 2)  # [B,H,T,d]
# k und v genauso
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
out = out.transpose(1, 2).contiguous().view(B, T, D)
out = self.out_proj(out)

Die Köpfe mischen Informationen separat, werden danach aber zusammengeführt. Eine Ausgabeprojektion kann ihre Ergebnisse kombinieren. Bei gleicher Gesamtbreite führt „mehr Köpfe“ nicht einfach zu einer proportional größeren Parameterzahl der Standardprojektionen. Stattdessen werden die einzelnen Köpfe schmaler.

Der Präfixtest

Erzeuge zwei Sequenzen, die bis Position 4 identisch sind und danach verschiedene Tokens besitzen. Bei einem kausalen Modell müssen die Logits bis Position 4 gleich bleiben, abgesehen von numerischem Rundungsrauschen. Setze das Modell in Eval-Modus und schalte Zufallsquellen wie Dropout aus.

python
model.eval()
a = torch.tensor([[1, 5, 6, 7, 8, 9]])
b = torch.tensor([[1, 5, 6, 7, 8, 3]])
with torch.no_grad():
    la, lb = model(a), model(b)
assert torch.allclose(la[:, :5], lb[:, :5], atol=1e-5)

Unser Testprojekt benutzt diese Idee. Er ist viel aussagekräftiger als ein visueller Blick auf einen Loss. Eine fehlende Maske kann den Trainings-Loss stark verbessern, gerade weil sie die Aufgabe unzulässig leicht macht.

Kausalität gilt auch für andere Schichten

LayerNorm oder RMSNorm normalisieren in unserem Modell nur die letzte Featureachse einer Position. Würdest du versehentlich über die Zeitachse normalisieren, könnten zukünftige Positionen in die Statistiken einfließen. Auch eine Datenvorverarbeitung, die Labels in den Input schreibt, kann die kausale Attention umgehen. Ein korrekter Attentionblock reicht daher nicht als Gesamtbeweis; teste das vollständige Modell.

Das Browserexperiment

Klicke auf eine Query-Zeile in der Matrix. Es zeigt die erlaubten Positionen, keine gelernten Beziehungen. Schalte die Maske ab und untersuche, welche zusätzlichen Informationen das Trainingsziel dann verraten würden. Die Matrix ist absichtlich als Strukturansicht beschriftet: Eine erlaubte Verbindung muss vom Modell nicht stark gewichtet werden.