Tokenwerk · Das LLM-Lehrbuch

Kapitel 19 · III · Den Transformer bauen · 9 Minuten

Der vollständige Transformer-Block

Attention allein genügt nicht. Residualpfad, Normalisierung und Feed-Forward-Netz arbeiten zusammen.

Die bekannten Teile werden zusammengesetzt

Ein Transformerblock ist eine wiederholbare Recheneinheit. Er erhält für jede Textposition eine Zahlenliste und gibt wieder eine gleich lange Zahlenliste aus. Er verbindet zwei bereits bekannte Ideen: Attention sammelt Informationen aus dem Kontext; ein kleines neuronales Netz verarbeitet die Zahlen an jeder Position weiter.

Wir bauen keinen neuen geheimnisvollen Mechanismus. Wir ordnen bekannte Rechenschritte und ergänzen zwei Hilfen fürs Training: direkte Additionswege und Normalisierung. Normalisierung bedeutet hier, die Größenordnung der Zahlen nach einer festgelegten Rechnung anzupassen.

Ein Ergebnis ergänzen, statt alles zu ersetzen

Stell dir eine Positionsdarstellung [2, 5] vor. Eine Operation berechnet die Änderung [0.3, -0.2]. Wir addieren komponentenweise und erhalten [2.3, 4.8]. Die ursprüngliche Darstellung bleibt als direkter Summand erhalten. Das ist eine Residualverbindung. Residual bedeutet hier: Die Operation liefert einen ergänzenden Beitrag.

In Kurzform schreiben wir y=x+f(x)y=x+f(x). x ist die Eingabe, f die Operation und y das Ergebnis. f(x) muss dieselbe Form wie x haben, damit die Addition möglich ist. Beim Rückwärtsrechnen teilt sich der Weg: Ein Teil geht direkt durch die Addition zurück, der andere durch f. Die Beiträge werden anschließend addiert. Erinnere dich an die verzweigten Rechenwege im Kapitel zur Backpropagation.

Das erleichtert häufig das Training tiefer Netze. Es garantiert aber nicht, dass jedes beliebig tiefe Netz mit jeder Lernrate stabil lernt.

Normalisieren mit zwei Zahlen

Nimm die Liste [1, 3]. Ihr Mittelwert ist (1 + 3) / 2 = 2. Ziehe ihn von beiden Werten ab: Es bleibt [-1, 1]. Die durchschnittliche quadrierte Abweichung ist (1 + 1) / 2 = 1. Diese Zahl heißt Varianz. Ihre Quadratwurzel heißt Standardabweichung; sie ist hier ebenfalls 1.

Teilen wir die Abweichungen durch die Standardabweichung, erhalten wir wieder [-1, 1]. Aus [10, 30] entstünde ohne Schutzterm dieselbe normalisierte Liste. Das ursprüngliche Niveau und die ursprüngliche Größenordnung werden damit entfernt.

Was passiert bei [2, 2]? Beide Abweichungen und die Varianz sind null. Wir dürfen nicht durch null teilen. Deshalb addieren wir vor dem Wurzelziehen eine sehr kleine positive Zahl, Epsilon. Sie dient als Schutzterm. Zusätzlich lernt das Modell für jede Komponente einen Multiplikator und eine Verschiebung. Es darf die standardisierten Werte damit wieder passend skalieren.

LayerNorm: die genaue Schreibweise

Die Normalisierung über die Komponenten einer einzelnen Position heißt LayerNorm. Unsere Implementation berechnet sie für jede Position separat:

LN⁡(x)i=γixi−μσ2+ϵ+βi.\operatorname{LN}(x)_i=\gamma_i\frac{x_i-\mu}{\sqrt{\sigma^2+\epsilon}}+\beta_i.
Zeichen Bedeutung in dieser Formel
x Zahlenliste einer Position
i Welche Komponente der Liste wir gerade berechnen
μ\mu (mü) Mittelwert aller Komponenten dieser Liste
σ2\sigma^2 (sigma zum Quadrat) Varianz: Mittelwert der quadrierten Abweichungen
ϵ\epsilon (epsilon) Kleine positive Schutzkonstante
γi\gamma_i (gamma) Trainierbarer Multiplikator für Komponente i
βi\beta_i (beta) Trainierbare Verschiebung für Komponente i
LN⁡(x)i\operatorname{LN}(x)_i Die fertig normalisierte i-te Komponente

Lies von innen nach außen: Mittelwert abziehen, durch die geschützte Standardabweichung teilen, mit Gamma multiplizieren, Beta addieren. Bei Gamma 1 und Beta 0 bleibt die standardisierte Liste erhalten. Epsilon führt zu einer kleinen Abweichung von der idealen Rechnung oben.

nn.LayerNorm(D) verarbeitet jeweils die letzten D Komponenten. Die verschiedenen Texte eines Batches und die verschiedenen Positionen werden dabei nicht zu einer gemeinsamen Statistik vermischt. LayerNorm begrenzt die Ausgabe auch nicht auf das Intervall von null bis eins.

Das kleine Netz innerhalb des Blocks

Nach der Attention folgt an jeder Position ein Feed-Forward-Netz, auch MLP genannt. MLP steht für Multilayer Perceptron: ein Netz aus mehreren aufeinanderfolgenden Schichten. Das Grundprinzip kennst du aus unserem ersten kleinen Netz.

Eine erste lineare Schicht erzeugt aus D Komponenten eine breitere Zwischenliste, hier mit 4D Komponenten. Eine Aktivierungsfunktion verändert die Werte nichtlinear. Eine zweite lineare Schicht bringt die Breite zurück auf D. Bei D = 8 heißt das: 8 → 32 → 8.

Wir verwenden GELU, eine glatte Aktivierungsfunktion. Anders als ReLU schneidet sie negative Werte nicht hart bei null ab; sie gewichtet sie weich. Ihre genaue Spezialform ist für das Zusammensetzen dieses Blocks nicht nötig. Entscheidend ist dieselbe Eigenschaft wie im Aktivierungskapitel: Die Gesamtoperation kann mehr ausdrücken als eine einzige lineare Schicht.

MLP⁡(x)=W2GELU⁡(W1x+b1)+b2.\operatorname{MLP}(x)=W_2\operatorname{GELU}(W_1x+b_1)+b_2.
Zeichen Bedeutung
W1W_1, b1b_1 Gewichtstabelle und Bias der ersten Schicht
W1x+b1W_1x+b_1 Breitere Zwischenliste, hier mit 4D Komponenten
GELU Aktivierungsfunktion, einzeln auf jede Komponente angewandt
W2W_2, b2b_2 Gewichtstabelle und Bias der zweiten Schicht
MLP(x) Ausgabe mit wieder D Komponenten

Wir schreiben hier mathematisch mit Spaltenvektoren. PyTorch speichert und multipliziert seine Batches entsprechend seiner eigenen Formkonvention; nn.Linear übernimmt die passende Anordnung.

python
self.mlp = nn.Sequential(
    nn.Linear(D, 4 * D),
    nn.GELU(),
    nn.Linear(4 * D, D),
)

Dasselbe MLP mit denselben Parametern wird auf jede Position angewendet. Es mischt selbst keine verschiedenen Positionen. Dafür ist die Attention zuständig.

Zwei Ergänzungen, zwei Normalisierungen

Jetzt lässt sich der Block in Worten lesen:

  1. Normalisiere die Eingabe und berechne daraus Attention.
  2. Addiere dieses Ergebnis zur ursprünglichen Eingabe. Nenne die neue Zwischenliste u.
  3. Normalisiere u und verarbeite sie im MLP.
  4. Addiere dessen Ergebnis zu u. Das ist die Ausgabe y.
u=x+Attention⁡(Norm⁡1(x)),u=x+\operatorname{Attention}(\operatorname{Norm}_1(x)),
y=u+MLP⁡(Norm⁡2(u)).y=u+\operatorname{MLP}(\operatorname{Norm}_2(u)).

x ist die Blockeingabe, u das Ergebnis nach der ersten Ergänzung und y die Blockausgabe. Norm₁ und Norm₂ sind zwei getrennte LayerNorm-Module mit eigenen trainierbaren Zahlen. Pre-Norm bedeutet: Die Normalisierung steht vor Attention beziehungsweise MLP. Der Originaltransformer von 2017 normalisierte erst nach der Addition (Post-Norm); Pre-Norm hat sich später als stabiler zu trainieren erwiesen (Xiong et al., 2020). Beim zweiten Additionsweg wird u addiert, nicht nochmals das alte x.

python
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x

Woher kommt die Positionsinformation?

Unsere Tokenembedding-Tabelle gibt demselben Token zunächst überall dieselbe Liste. Für Sprache muss aber auch die Stelle in der Folge berücksichtigt werden. Im klassischen GPT-2-artigen Decoder addieren wir eine zweite gelernte Liste, das Positionsembedding. Der ursprüngliche Transformer von 2017 verwendete stattdessen feste Sinus-Kosinus-Kodierungen.

Für eine Tokenliste [0.2, 0.5] und eine Positionsliste [0.1, -0.2] erhält der erste Block [0.3, 0.3]. Beide Tabellen werden trainiert. Der Index t einer Position wählt aus der Positionstabelle die entsprechende Zeile aus.

python
pos = torch.arange(T, device=ids.device)
x = self.token_embedding(ids) + self.position_embedding(pos)

T ist die Länge der aktuellen Tokenfolge. arange(T) erzeugt die Positionsnummern 0 bis T−1. device sorgt dafür, dass beide Zahlenfelder auf demselben Rechengerät liegen.

Die gelernte Positionstabelle hat eine feste Anzahl Zeilen. Bei Kontextlänge 256 gibt es keine Zeile für Position 10000. Unser späterer Generator behält daher während der Erzeugung höchstens die letzten context Tokens (zum Beispiel 256) und nummeriert dieses Fenster neu. Alte, entfernte Tokens stehen dem Modell bei diesem Rechenschritt nicht mehr zur Verfügung.

Der vollständige Decoder

Der Ablauf lautet: Token- und Positionsembeddings addieren, mehrere Transformerblöcke anwenden, noch einmal normalisieren, schließlich eine lineare Ausgabeschicht ausführen. Diese letzte Schicht heißt Ausgabekopf. Sie produziert pro Position einen Logit für jeden möglichen nächsten Token. Bei V = 100 Vokabulareinträgen sind das 100 Zahlen.

Der Kopf wählt noch keinen Token aus. Softmax und die spätere Auswahlstrategie machen daraus eine Fortsetzung. Dieselben gelernten Parameter werden für alle Positionen und für jeden neuen Generierungsschritt benutzt.

Vertiefung: Parameter grob zählen

Eine D-mal-D-Gewichtstabelle enthält D² Zahlen. Vier solche Tabellen für Q, K, V und die Attention-Ausgabe enthalten ungefähr 4D² Gewichte. Die beiden MLP-Tabellen haben D mal 4D und 4D mal D Einträge, zusammen 8D². Ein Block enthält damit ungefähr 12D² Gewichte. Biases und Normparameter kommen hinzu.

Für L Blöcke, V Vokabulareinträge und T gelernte Positionen ergibt sich bei geteilter Ein-/Ausgabetabelle die Näherung:

N≈VD+12LD2+TD.N\approx VD+12LD^2+TD.

N ist die Gesamtzahl der Parameter; D die Breite einer Positionsliste; L die Blockzahl; V die Vokabulargröße; T die maximale Positionszahl. Die drei Summanden zählen Tokenembedding, Blöcke und Positionstabelle. Bei V = 100, D = 8, L = 2 und T = 16 sind das 800 + 1536 + 128 = 2464 Gewichte zuzüglich der ausgelassenen kleinen Parametergruppen.

Diese Schätzung gilt für den hier beschriebenen klassischen Block. Spätere Architekturänderungen verändern die Rechnung. Der Code kann die tatsächlich vorhandenen Parameter direkt zählen: sum(p.numel() for p in model.parameters()). numel() liefert die Anzahl der Zahlen eines Parameterfelds.