Was wir bauen
Dein GPT entspricht ungefähr GPT-2 von 2019. In Kapitel 25 hast du drei Bauteile kennengelernt, die heute in fast allen großen Modellen stecken: RMSNorm, RoPE und SwiGLU. Jetzt baust du sie ein – und misst, ob sie bei unseren Märchen tatsächlich helfen.
Damit du beide Varianten vergleichen kannst, bekommt Config einen Schalter: modern=False ist dein bisheriges Modell, modern=True die moderne Version. Alles andere bleibt gleich.
Der Schalter
@dataclassRMSNorm
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weightWie LayerNorm, nur ohne den Durchschnitt abzuziehen (Kapitel 25): Quadrate mitteln, Wurzel ziehen, durch diese Größe teilen, mit einem lernbaren Gewicht malnehmen. torch.rsqrt heißt „1 geteilt durch die Wurzel“ – eine Rechnung statt zwei. eps ist die Schutzzahl gegen das Teilen durch null.
RoPE: Position durch Drehen
def rope(x):
"""Dreht Zahlenpaare je nach Position (x: [B, H, T, d])."""
T, d = x.shape[-2], x.shape[-1]
speed = 10000 ** (-torch.arange(0, d, 2, device=x.device) / d) # jedes Paar dreht anders schnell
angle = torch.arange(T, device=x.device)[:, None] * speed[None, :] # [T, d/2]
cos, sin = angle.cos(), angle.sin()
even, odd = x[..., 0::2], x[..., 1::2]
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)Das ist die Uhrzeiger-Idee aus Kapitel 25 in sieben Zeilen:
speedlegt für jedes Zahlenpaar fest, wie schnell es sich dreht – das erste Paar wie ein Sekundenzeiger, das letzte wie ein Stundenzeiger.angleist Position mal Geschwindigkeit: An Position 0 wird nicht gedreht, an Position 10 zehnmal so weit wie an Position 1.evenundoddsind die beiden Zahlen jedes Paars. Die Zeile mitcosundsinist die Drehformel aus der Mathe-Ecke von Kapitel 25.stackundflattenlegen die gedrehten Paare wieder in der ursprünglichen Reihenfolge ab.
Gedreht werden nur Query und Key, nicht Value. Deshalb braucht die Attention genau zwei neue Zeilen:
class SelfAttention(nn.Module):
def __init__(self, cfg):
super().__init__()
self.heads = cfg.heads
self.modern = cfg.modern
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)
if self.modern:
q, k = rope(q), rope(k) # Position durch Drehen statt Tabelle
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))Und weil die Position jetzt in der Drehung steckt, braucht das moderne GPT keine Positionstabelle mehr.
SwiGLU: das MLP mit Türsteher
class MLP(nn.Module):
def __init__(self, cfg):
super().__init__()
self.modern = cfg.modern
hidden = 4 * cfg.dim if not cfg.modern else 8 * cfg.dim // 3
self.up = nn.Linear(cfg.dim, hidden)
if cfg.modern:
self.gate = nn.Linear(cfg.dim, hidden) # SwiGLU: der Türsteher
self.down = nn.Linear(hidden, cfg.dim)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x):
if self.modern:
return self.drop(self.down(F.silu(self.gate(x)) * self.up(x)))
return self.drop(self.down(F.gelu(self.up(x)))) # breiter, Knick, wieder schmalIm modernen Modus gibt es einen zweiten Weg: gate, der Türsteher. F.silu(self.gate(x)) entscheidet Stelle für Stelle, wie viel von self.up(x) durchkommt. Weil es jetzt drei Gewichtstabellen statt zwei sind, machen wir die Zwischenbreite kleiner: 8D/3 statt 4D. So bleibt die Parameterzahl vergleichbar.
Block und GPT
class Block(nn.Module):
def __init__(self, cfg):
super().__init__()
Norm = RMSNorm if cfg.modern else nn.LayerNorm
self.norm1 = Norm(cfg.dim)
self.attn = SelfAttention(cfg)
self.norm2 = Norm(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 xclass GPT(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.tokens = nn.Embedding(cfg.vocab_size, cfg.dim) # was
self.positions = None if cfg.modern else nn.Embedding(cfg.context, cfg.dim) # wo (modern: RoPE)
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
self.norm = RMSNorm(cfg.dim) if cfg.modern else 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)
if self.positions is not None:
x = x + 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]Im Block entscheidet der Schalter nur, welche Normalisierung benutzt wird. Im GPT fällt bei modern=True die Positionstabelle weg (positions = None), und in forward wird sie nur addiert, wenn es sie gibt.
Die ganze Datei
Das ist die finale Fassung von model.py. Sie kann beides: klassisch und modern.
# 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)
modern: bool = False # RMSNorm + RoPE + SwiGLU statt LayerNorm + Positionstabelle + GELU (Kapitel 25)
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight
def rope(x):
"""Dreht Zahlenpaare je nach Position (x: [B, H, T, d])."""
T, d = x.shape[-2], x.shape[-1]
speed = 10000 ** (-torch.arange(0, d, 2, device=x.device) / d) # jedes Paar dreht anders schnell
angle = torch.arange(T, device=x.device)[:, None] * speed[None, :] # [T, d/2]
cos, sin = angle.cos(), angle.sin()
even, odd = x[..., 0::2], x[..., 1::2]
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)
class SelfAttention(nn.Module):
def __init__(self, cfg):
super().__init__()
self.heads = cfg.heads
self.modern = cfg.modern
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)
if self.modern:
q, k = rope(q), rope(k) # Position durch Drehen statt Tabelle
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.modern = cfg.modern
hidden = 4 * cfg.dim if not cfg.modern else 8 * cfg.dim // 3
self.up = nn.Linear(cfg.dim, hidden)
if cfg.modern:
self.gate = nn.Linear(cfg.dim, hidden) # SwiGLU: der Türsteher
self.down = nn.Linear(hidden, cfg.dim)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x):
if self.modern:
return self.drop(self.down(F.silu(self.gate(x)) * self.up(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__()
Norm = RMSNorm if cfg.modern else nn.LayerNorm
self.norm1 = Norm(cfg.dim)
self.attn = SelfAttention(cfg)
self.norm2 = Norm(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 = None if cfg.modern else nn.Embedding(cfg.context, cfg.dim) # wo (modern: RoPE)
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
self.norm = RMSNorm(cfg.dim) if cfg.modern else 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)
if self.positions is not None:
x = x + 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]Ausführen und vergleichen
Prüf zuerst, dass das klassische Modell noch genau gleich funktioniert:
python check.pyDie Ausgabe muss dieselbe sein wie in Schritt 4. Dann schaltest du in train.py die moderne Variante ein. Ändere eine einzige Zeile:
cfg = Config(vocab_size=tok.vocab_size, modern=True)python train.py923,432 Parameter auf cpu
...
Bester Validierungs-Loss 3.383, gespeichert in data/model.ptDer direkte Vergleich auf denselben Daten, mit denselben Einstellungen:
| Klassisch | Modern | |
|---|---|---|
| Parameter | 940.800 | 923.432 |
| Bester Validierungs-Loss | 3,535 | 3,383 |
| Trainingszeit (bei mir) | 114 s | 126 s |
Das moderne Modell ist besser – und das mit weniger Parametern. Ein Loss-Unterschied von 0,15 klingt klein, ist aber deutlich: Das Modell gibt dem richtigen nächsten Token im Schnitt rund 16 Prozent mehr Wahrscheinlichkeit (e hoch 0,15 ≈ 1,16). Dafür dauert jeder Schritt etwas länger, vor allem wegen der Drehungen.
Genau so arbeitet man in der Forschung: Hypothese („moderne Bauteile helfen“), eine Sache ändern, messen, ehrlich vergleichen. Wenn du es genau wissen willst, schalte die drei Bauteile einzeln ein – das wäre eine echte Ablation aus Kapitel 25.
Nach dem Training mit modern=True passen generate.py, sft.py und chat.py ohne Änderung: Sie lesen die Einstellungen aus dem gespeicherten Modell.
Wie es weitergehen kann
Dein Modell hat jetzt dieselben Bauteile wie viele aktuelle Open-Source-Modelle. Was ihm noch fehlt, sind vor allem Größe und Daten. Ideen für die nächsten Schritte stehen im folgenden Kapitel, zum Beispiel:
- GQA aus Kapitel 26: Mehrere Query-Köpfe teilen sich Keys und Values.
- Ein KV-Cache für schnelleres Generieren.
- Mehr Daten und ein größeres Modell, am besten auf einer Grafikkarte.
Kurz gemerkt
- Ein Schalter
moderntauscht drei Bauteile: RMSNorm statt LayerNorm, RoPE statt Positionstabelle, SwiGLU statt GELU-MLP. - RoPE dreht nur Query und Key. Dafür braucht die Attention zwei Zeilen mehr, und die Positionstabelle fällt weg.
- Auf unseren Märchen ist das moderne Modell messbar besser – mit weniger Parametern.