Tokenwerk · Das LLM-Lehrbuch

Kapitel 36 · VII · Bau dein LLM · 12 Minuten

Schritt 4: Das GPT-Modell

Attention mit Maske und Köpfen, MLP, Blöcke, Embeddings: Das komplette GPT in rund 90 Zeilen – und drei Prüfungen, bevor eine Sekunde trainiert wird.

Was wir bauen

Jetzt kommt das Herzstück: ein echtes GPT. Es ist genau der Decoder aus Kapitel 19 – Embeddings, Attention, MLP, Residualverbindungen, Normalisierung –, nur diesmal als Code. Alles landet in der Datei model.py, gut 90 Zeilen.

Der Bauplan von innen nach außen:

  1. Config – alle Größen an einem Ort.
  2. SelfAttention – Positionen tauschen sich aus (Kapitel 17 und 18).
  3. MLP – jede Position rechnet für sich weiter.
  4. Block – Attention und MLP mit Normalisierung und Residual zusammengesteckt.
  5. GPT – Embeddings, ein Stapel Blöcke, Ausgabekopf.

Am Ende prüfen wir mit check.py, ob alles richtig verdrahtet ist – bevor wir eine einzige Sekunde trainieren.

Die Einstellungen

python
# Ein GPT von Grund auf: Embedding, Attention, MLP, Blöcke (Kapitel 15–19)
import math
from dataclasses import dataclass

import torch
import torch.nn as nn
import torch.nn.functional as F
python
@dataclass

Eine dataclass ist eine Klasse, die nur Werte aufbewahrt. So steht jede Größe genau einmal da, und alle Bausteine lesen sie aus cfg. Die Buchstaben kennst du aus Kapitel 11: T, D, H und L. Mit diesen Werten bekommt unser Modell knapp eine Million Parameter – winzig gegen ChatGPT, aber genug, um Märchen zu lernen.

dropout ist neu: Beim Training wird zufällig ein Fünftel der Werte auf null gesetzt. Das zwingt das Modell, sich nicht auf einzelne Verbindungen zu verlassen, und bremst das Auswendiglernen. Warum wir das bei unserem kleinen Märchentext brauchen, zeigt Schritt 5.

Self-Attention

python
class SelfAttention(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.heads = cfg.heads
        self.qkv = nn.Linear(cfg.dim, 3 * cfg.dim)               # Q, K und V in einem Rutsch
        self.out = nn.Linear(cfg.dim, cfg.dim)                   # mischt die Köpfe wieder zusammen
        self.drop = nn.Dropout(cfg.dropout)
        mask = torch.tril(torch.ones(cfg.context, cfg.context, dtype=torch.bool))
        self.register_buffer("mask", mask)                       # kausale Maske, kein Parameter

    def forward(self, x):
        B, T, D = x.shape
        d = D // self.heads
        q, k, v = self.qkv(x).split(D, dim=-1)                   # je [B, T, D]
        # in Köpfe aufteilen: [B, T, D] -> [B, H, T, d]
        q = q.view(B, T, self.heads, d).transpose(1, 2)
        k = k.view(B, T, self.heads, d).transpose(1, 2)
        v = v.view(B, T, self.heads, d).transpose(1, 2)
        scores = q @ k.transpose(-2, -1) / math.sqrt(d)          # vergleichen: [B, H, T, T]
        scores = scores.masked_fill(~self.mask[:T, :T], float("-inf"))   # Zukunft sperren
        weights = scores.softmax(dim=-1)                         # Prozente pro Zeile
        y = weights @ v                                          # mischen: [B, H, T, d]
        y = y.transpose(1, 2).contiguous().view(B, T, D)         # Köpfe wieder nebeneinander
        return self.drop(self.out(y))

Das ist Kapitel 17 und 18 in 25 Zeilen. Gehen wir forward Zeile für Zeile durch:

  • Q, K, V berechnen: self.qkv ist eine einzige lineare Schicht, die aus jeder Zahlenliste der Breite D direkt drei Listen macht – Query, Key und Value. Das ist dasselbe wie drei getrennte Schichten, nur schneller. split(D, dim=-1) schneidet das Ergebnis wieder in drei Teile.
  • In Köpfe aufteilen: view(B, T, H, d) schneidet jede Liste in H Stücke der Breite d = D/H. transpose(1, 2) stellt die Kopf-Achse nach vorne: aus [B, T, H, d] wird [B, H, T, d]. Jetzt rechnet jeder Kopf für sich (Kapitel 18).
  • Vergleichen: q @ k.transpose(-2, -1) berechnet alle Skalarprodukte zwischen Queries und Keys auf einmal – die Scorematrix der Form [B, H, T, T]. Geteilt durch √d, wie in Kapitel 17.
  • Zukunft sperren: self.mask ist ein Dreieck aus True/False. Überall, wo False steht – also rechts oberhalb der Diagonale –, setzt masked_fill den Score auf −∞. register_buffer heißt: Die Maske gehört zum Modell, wird aber nicht trainiert.
  • Gewichten und mischen: softmax macht aus jeder Zeile Prozente, weights @ v mischt die Values.
  • Köpfe zusammenlegen: Rücktransponieren, alle Köpfe wieder nebeneinander, dann mischt self.out die Ergebnisse.

Das MLP

python
class MLP(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.up = nn.Linear(cfg.dim, 4 * cfg.dim)
        self.down = nn.Linear(4 * cfg.dim, cfg.dim)
        self.drop = nn.Dropout(cfg.dropout)

    def forward(self, x):
        return self.drop(self.down(F.gelu(self.up(x))))          # breiter, Knick, wieder schmal

Breiter machen (D → 4D), Knick (GELU), wieder schmal (4D → D). Exakt wie in Kapitel 19.

Der Block

python
class Block(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.norm1 = nn.LayerNorm(cfg.dim)
        self.attn = SelfAttention(cfg)
        self.norm2 = nn.LayerNorm(cfg.dim)
        self.mlp = MLP(cfg)

    def forward(self, x):
        x = x + self.attn(self.norm1(x))                         # Positionen tauschen sich aus
        x = x + self.mlp(self.norm2(x))                          # jede Position rechnet für sich
        return x

Zwei Zeilen in forward, und doch der ganze Transformer-Block: Normalisieren, Attention, Ergebnis dazuaddieren. Dann Normalisieren, MLP, wieder dazuaddieren. Das x = x + ... ist die Residualverbindung – ergänzen statt ersetzen.

Das ganze GPT

python
class GPT(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.tokens = nn.Embedding(cfg.vocab_size, cfg.dim)      # was
        self.positions = nn.Embedding(cfg.context, cfg.dim)      # wo
        self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
        self.norm = nn.LayerNorm(cfg.dim)
        self.head = nn.Linear(cfg.dim, cfg.vocab_size, bias=False)
        self.head.weight = self.tokens.weight                    # Weight Tying: eine Tabelle, zwei Jobs
        self.apply(self.init_weights)

    @staticmethod
    def init_weights(module):
        # kleine Zufallszahlen als Start (Kapitel 20) – sonst sind die ersten Logits riesig
        if isinstance(module, (nn.Linear, nn.Embedding)):
            nn.init.normal_(module.weight, std=0.02)
        if isinstance(module, nn.Linear) and module.bias is not None:
            nn.init.zeros_(module.bias)

    def forward(self, ids):
        B, T = ids.shape
        x = self.tokens(ids) + self.positions(torch.arange(T, device=ids.device))
        for block in self.blocks:
            x = block(x)
        return self.head(self.norm(x))                           # Logits: [B, T, V]
  • tokens und positions sind zwei Embedding-Tabellen: eine für was (welches Token), eine für wo (welche Position). Beide Listen werden addiert.
  • nn.ModuleList hält den Stapel von Blöcken. Eine normale Python-Liste ginge nicht, denn dann fände PyTorch die Gewichte darin nicht (Kapitel 20).
  • self.head.weight = self.tokens.weight ist Weight Tying aus Kapitel 15: Eingangs- und Ausgangstabelle teilen sich dieselben Zahlen.
  • init_weights gibt allen Gewichten kleine Startwerte. Das ist kein Detail! Ohne diese Zeile startet das Training bei einem Loss von 87 statt 6,9. Grund: Die Embedding-Tabelle wird von PyTorch mit großen Zufallszahlen gefüllt, und durchs Weight Tying landen die auch im Ausgabekopf. Riesige Logits bedeuten ein Modell, das sich sehr sicher ist – und sehr falsch.
  • forward erzählt die ganze Geschichte aus Kapitel 19 in fünf Zeilen: nachschlagen, Position addieren, durch alle Blöcke, normalisieren, Ausgabekopf.

Die ganze Datei

python
# Ein GPT von Grund auf: Embedding, Attention, MLP, Blöcke (Kapitel 15–19)
import math
from dataclasses import dataclass

import torch
import torch.nn as nn
import torch.nn.functional as F


@dataclass
class Config:
    vocab_size: int
    context: int = 128      # T: wie viele Tokens das Modell auf einmal sieht
    dim: int = 128          # D: Zahlen pro Token
    heads: int = 4          # H: Attention-Köpfe
    layers: int = 4         # L: Transformer-Blöcke
    dropout: float = 0.2    # gegen Auswendiglernen (Kapitel 23)


class SelfAttention(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.heads = cfg.heads
        self.qkv = nn.Linear(cfg.dim, 3 * cfg.dim)               # Q, K und V in einem Rutsch
        self.out = nn.Linear(cfg.dim, cfg.dim)                   # mischt die Köpfe wieder zusammen
        self.drop = nn.Dropout(cfg.dropout)
        mask = torch.tril(torch.ones(cfg.context, cfg.context, dtype=torch.bool))
        self.register_buffer("mask", mask)                       # kausale Maske, kein Parameter

    def forward(self, x):
        B, T, D = x.shape
        d = D // self.heads
        q, k, v = self.qkv(x).split(D, dim=-1)                   # je [B, T, D]
        # in Köpfe aufteilen: [B, T, D] -> [B, H, T, d]
        q = q.view(B, T, self.heads, d).transpose(1, 2)
        k = k.view(B, T, self.heads, d).transpose(1, 2)
        v = v.view(B, T, self.heads, d).transpose(1, 2)
        scores = q @ k.transpose(-2, -1) / math.sqrt(d)          # vergleichen: [B, H, T, T]
        scores = scores.masked_fill(~self.mask[:T, :T], float("-inf"))   # Zukunft sperren
        weights = scores.softmax(dim=-1)                         # Prozente pro Zeile
        y = weights @ v                                          # mischen: [B, H, T, d]
        y = y.transpose(1, 2).contiguous().view(B, T, D)         # Köpfe wieder nebeneinander
        return self.drop(self.out(y))


class MLP(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.up = nn.Linear(cfg.dim, 4 * cfg.dim)
        self.down = nn.Linear(4 * cfg.dim, cfg.dim)
        self.drop = nn.Dropout(cfg.dropout)

    def forward(self, x):
        return self.drop(self.down(F.gelu(self.up(x))))          # breiter, Knick, wieder schmal


class Block(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.norm1 = nn.LayerNorm(cfg.dim)
        self.attn = SelfAttention(cfg)
        self.norm2 = nn.LayerNorm(cfg.dim)
        self.mlp = MLP(cfg)

    def forward(self, x):
        x = x + self.attn(self.norm1(x))                         # Positionen tauschen sich aus
        x = x + self.mlp(self.norm2(x))                          # jede Position rechnet für sich
        return x


class GPT(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.tokens = nn.Embedding(cfg.vocab_size, cfg.dim)      # was
        self.positions = nn.Embedding(cfg.context, cfg.dim)      # wo
        self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
        self.norm = nn.LayerNorm(cfg.dim)
        self.head = nn.Linear(cfg.dim, cfg.vocab_size, bias=False)
        self.head.weight = self.tokens.weight                    # Weight Tying: eine Tabelle, zwei Jobs
        self.apply(self.init_weights)

    @staticmethod
    def init_weights(module):
        # kleine Zufallszahlen als Start (Kapitel 20) – sonst sind die ersten Logits riesig
        if isinstance(module, (nn.Linear, nn.Embedding)):
            nn.init.normal_(module.weight, std=0.02)
        if isinstance(module, nn.Linear) and module.bias is not None:
            nn.init.zeros_(module.bias)

    def forward(self, ids):
        B, T = ids.shape
        x = self.tokens(ids) + self.positions(torch.arange(T, device=ids.device))
        for block in self.blocks:
            x = block(x)
        return self.head(self.norm(x))                           # Logits: [B, T, V]

Prüfen, bevor wir trainieren

Ein Fehler im Modell fällt beim Training oft erst nach Stunden auf – oder gar nicht, weil der Loss trotzdem sinkt. Deshalb prüfen wir vorher drei Dinge (Kapitel 18 und 20). Lege die Datei check.py an:

python
# Prüfungen, bevor wir trainieren (Kapitel 18 und 20)
import torch
import torch.nn.functional as F

from model import GPT, Config

torch.manual_seed(0)
cfg = Config(vocab_size=1024, dropout=0.0)
model = GPT(cfg).eval()

ids = torch.randint(0, cfg.vocab_size, (2, 16))
logits = model(ids)
print("Form der Ausgabe:", tuple(logits.shape), "(erwartet: (2, 16, 1024))")
print("Parameter:", f"{sum(p.numel() for p in model.parameters()):,}")

# 1. Unsere Attention muss dasselbe liefern wie PyTorchs eingebaute Version
attn = model.blocks[0].attn
x = torch.randn(2, 16, cfg.dim)
q, k, v = attn.qkv(x).split(cfg.dim, dim=-1)
split = lambda t: t.view(2, 16, cfg.heads, -1).transpose(1, 2)
reference = F.scaled_dot_product_attention(split(q), split(k), split(v), is_causal=True)
reference = attn.out(reference.transpose(1, 2).reshape(2, 16, cfg.dim))
print("Gleich wie PyTorch:", torch.allclose(attn(x), reference, atol=1e-5))

# 2. Spick-Test: Ein anderes Ende darf frühere Vorhersagen nicht verändern
a = torch.tensor([[1, 5, 6, 7, 8, 9]])
b = torch.tensor([[1, 5, 6, 7, 8, 3]])
same = torch.allclose(model(a)[:, :5], model(b)[:, :5], atol=1e-5)
print("Schaut nicht in die Zukunft:", same)

# 3. Startfehler sollte ungefähr ln(1024) = 6.93 sein
targets = torch.randint(0, cfg.vocab_size, (2, 16))
loss = F.cross_entropy(logits.view(-1, cfg.vocab_size), targets.view(-1))
print(f"Startfehler: {loss.item():.2f}")
  1. Stimmt unsere Attention? PyTorch hat eine eingebaute, stark optimierte Attention-Funktion. Wenn unsere eigene Rechnung dasselbe Ergebnis liefert, haben wir keinen Rechenfehler gemacht.
  2. Schaut das Modell in die Zukunft? Der Spick-Test aus Kapitel 18: Zwei Folgen, die erst ganz hinten verschieden sind, müssen vorher dieselben Vorhersagen haben.
  3. Startet das Modell neutral? Ein untrainiertes Modell sollte keine Meinung haben. Dann ist der Fehler ungefähr ln(1024) ≈ 6,93 – so viel, als würde es zwischen allen 1024 Tokens gleichmäßig raten.

.eval() und dropout=0.0 schalten den Zufall des Dropouts aus, damit die Vergleiche exakt sind.

bash
python check.py
text
Form der Ausgabe: (2, 16, 1024) (erwartet: (2, 16, 1024))
Parameter: 940,800
Gleich wie PyTorch: True
Schaut nicht in die Zukunft: True
Startfehler: 6.93

Drei Mal grünes Licht. Dein GPT ist richtig gebaut.

Wenn etwas schiefgeht

  • Gleich wie PyTorch: False – Meist ein Fehler in der Maske oder beim Teilen durch √d. Vergleiche die Zeilen mit scores genau mit dem Buch.
  • Schaut nicht in die Zukunft: False – Die Maske fehlt oder ist falsch herum. torch.tril (unteres Dreieck) ist richtig, torch.triu wäre falsch.
  • Startfehler deutlich über 7 – self.apply(self.init_weights) fehlt oder steht vor dem Weight Tying an der falschen Stelle.
  • RuntimeError: view size is not compatible – .contiguous() vor dem letzten view vergessen.

Kurz gemerkt

  • model.py baut Kapitel 15–19 nach: Embeddings, Self-Attention mit Maske und Köpfen, MLP, Block mit Residual, Ausgabekopf mit Weight Tying.
  • Kleine Startgewichte sind wichtig: Sonst beginnt das Training mit einem riesigen Fehler.
  • check.py prüft vor dem Training: Attention rechnet richtig, keine Blicke in die Zukunft, neutraler Start bei 6,93.