Tokenwerk · Das LLM-Lehrbuch

Kapitel 40 · VII · Bau dein LLM · 12 Minuten

Schritt 8: Modernisieren

RMSNorm, RoPE und SwiGLU in dein Modell einbauen – und messen, ob sie wirklich helfen. Spoiler: ja, mit weniger Parametern.

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

python
@dataclass

RMSNorm

python
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

Wie 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

python
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:

  • speed legt für jedes Zahlenpaar fest, wie schnell es sich dreht – das erste Paar wie ein Sekundenzeiger, das letzte wie ein Stundenzeiger.
  • angle ist Position mal Geschwindigkeit: An Position 0 wird nicht gedreht, an Position 10 zehnmal so weit wie an Position 1.
  • even und odd sind die beiden Zahlen jedes Paars. Die Zeile mit cos und sin ist die Drehformel aus der Mathe-Ecke von Kapitel 25.
  • stack und flatten legen 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:

python
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

python
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

Im 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

python
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
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 = 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.

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)
    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:

bash
python check.py

Die Ausgabe muss dieselbe sein wie in Schritt 4. Dann schaltest du in train.py die moderne Variante ein. Ändere eine einzige Zeile:

python
cfg = Config(vocab_size=tok.vocab_size, modern=True)
bash
python train.py
text
923,432 Parameter auf cpu
...
Bester Validierungs-Loss 3.383, gespeichert in data/model.pt

Der 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 modern tauscht 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.