Tokenwerk · Das LLM-Lehrbuch

Kapitel 37 · VII · Bau dein LLM · 9 Minuten

Schritt 5: Trainieren

Die Trainingsschleife mit allem, was dazugehört: AdamW, Warmup, Clipping, Validierung. Nach zwei Minuten hast du ein trainiertes Sprachmodell.

Was wir bauen

Das Modell steht, jetzt soll es lernen. Die Datei train.py ist die Trainingsschleife aus Kapitel 22: Batches ziehen, vorhersagen, Fehler messen, zurückrechnen, Gewichte anpassen – 2000-mal. Dazu kommen die Profi-Zutaten aus demselben Kapitel: AdamW, Warmup und Abkühlen der Lernrate, Gradient Clipping, regelmäßige Validierung und das Speichern des besten Stands.

Auf einem normalen Laptop dauert das ein bis drei Minuten. Danach hast du ein trainiertes Sprachmodell in data/model.pt.

Vorbereitung

python
# Die Trainingsschleife (Kapitel 22)
import math
import time

import numpy as np
import torch
import torch.nn.functional as F

from model import GPT, Config
from tokenizer import Tokenizer

STEPS, BATCH, LR, WARMUP = 2000, 32, 2e-3, 100
torch.manual_seed(42)
device = "cuda" if torch.cuda.is_available() else "cpu"

tok = Tokenizer.load("data/tokenizer.json")
train = np.load("data/train.npy")
valid = np.load("data/valid.npy")
cfg = Config(vocab_size=tok.vocab_size)
model = GPT(cfg).to(device)
print(f"{sum(p.numel() for p in model.parameters()):,} Parameter auf {device}")
  • Die vier großen Stellschrauben stehen ganz oben: 2000 Schritte, 32 Texte pro Batch, Lernrate 0,002, 100 Schritte Warmup.
  • torch.manual_seed(42) macht den Zufall wiederholbar: Gleicher Seed, gleicher Lauf.
  • Hast du eine NVIDIA-Grafikkarte, nimmt das Skript automatisch cuda. Sonst die CPU – das reicht völlig.

Batches ziehen

python
def get_batch(data):
    starts = np.random.randint(0, len(data) - cfg.context - 1, BATCH)
    x = np.stack([data[s : s + cfg.context] for s in starts])           # Eingabe
    y = np.stack([data[s + 1 : s + 1 + cfg.context] for s in starts])   # Ziel: um eins verschoben
    return torch.from_numpy(x).long().to(device), torch.from_numpy(y).long().to(device)

Genau wie beim Bigram-Modell, nur mit Fenstern der Länge cfg.context = 128: Zufällige Startstellen, ab jeder 128 Tokens als Eingabe und dieselben 128 Tokens um eins verschoben als Ziel. Jedes Fenster enthält also 128 Übungsaufgaben auf einmal.

Fehler messen – auch ehrlich

python
def loss_on(x, y):
    logits = model(x)
    return F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1))


@torch.no_grad()
def validation_loss():
    model.eval()
    losses = [loss_on(*get_batch(valid)).item() for _ in range(20)]
    model.train()
    return sum(losses) / len(losses)

loss_on ist der Fehler für einen Batch. validation_loss misst den Fehler auf den Validierungsdaten, die das Modell nie zum Lernen bekommt (Kapitel 21 und 23):

  • @torch.no_grad() sagt PyTorch: Hier wird nicht gelernt, also merk dir keinen Rechenweg. Das spart Zeit und Speicher.
  • model.eval() schaltet Dropout aus, model.train() danach wieder ein.
  • Wir mitteln über 20 zufällige Batches. Ein einzelner wäre zu zufällig.

Die Lernrate planen

python
def learning_rate(step):
    if step < WARMUP:
        return LR * (step + 1) / WARMUP                                  # langsam aufwärmen
    progress = (step - WARMUP) / (STEPS - WARMUP)
    return LR * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress)))  # sanft abkühlen

Warmup und Abkühlen aus Kapitel 22: In den ersten 100 Schritten steigt die Lernrate langsam auf 0,002. Danach sinkt sie entlang einer Kosinuskurve auf ein Zehntel. Große Schritte am Anfang, vorsichtige am Ende.

Die Schleife

python
optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=0.1)
start, best = time.time(), float("inf")
for step in range(STEPS + 1):
    for group in optimizer.param_groups:
        group["lr"] = learning_rate(step)
    x, y = get_batch(train)
    loss = loss_on(x, y)
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)              # Notbremse
    optimizer.step()
    if step % 200 == 0:
        val = validation_loss()
        print(f"Schritt {step:5d}  Train {loss.item():.3f}  Valid {val:.3f}  ({time.time() - start:.0f} s)")
        if val < best:                                                   # nur den besten Stand behalten
            best = val
            torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/model.pt")

print(f"Bester Validierungs-Loss {best:.3f}, gespeichert in data/model.pt")

Das Herz des Ganzen – und du erkennst jede Zeile wieder:

  • AdamW passt die Schrittgröße für jedes Gewicht einzeln an. weight_decay=0.1 zieht alle Gewichte sanft Richtung null.
  • Vor jedem Schritt setzen wir die geplante Lernrate.
  • Die fünf Kernzeilen: Batch, Loss, zero_grad, backward, step.
  • clip_grad_norm_ ist die Notbremse gegen riesige Steigungen.
  • Alle 200 Schritte messen wir den Validierungs-Loss. Ist er besser als je zuvor, speichern wir das Modell. So bleibt am Ende automatisch der beste Stand übrig, nicht der letzte.

Gespeichert werden die Einstellungen (cfg.__dict__) und alle Gewichte (state_dict()). Mehr braucht es nicht, um das Modell später wieder aufzubauen.

Die ganze Datei

python
# Die Trainingsschleife (Kapitel 22)
import math
import time

import numpy as np
import torch
import torch.nn.functional as F

from model import GPT, Config
from tokenizer import Tokenizer

STEPS, BATCH, LR, WARMUP = 2000, 32, 2e-3, 100
torch.manual_seed(42)
device = "cuda" if torch.cuda.is_available() else "cpu"

tok = Tokenizer.load("data/tokenizer.json")
train = np.load("data/train.npy")
valid = np.load("data/valid.npy")
cfg = Config(vocab_size=tok.vocab_size)
model = GPT(cfg).to(device)
print(f"{sum(p.numel() for p in model.parameters()):,} Parameter auf {device}")


def get_batch(data):
    starts = np.random.randint(0, len(data) - cfg.context - 1, BATCH)
    x = np.stack([data[s : s + cfg.context] for s in starts])           # Eingabe
    y = np.stack([data[s + 1 : s + 1 + cfg.context] for s in starts])   # Ziel: um eins verschoben
    return torch.from_numpy(x).long().to(device), torch.from_numpy(y).long().to(device)


def loss_on(x, y):
    logits = model(x)
    return F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1))


@torch.no_grad()
def validation_loss():
    model.eval()
    losses = [loss_on(*get_batch(valid)).item() for _ in range(20)]
    model.train()
    return sum(losses) / len(losses)


def learning_rate(step):
    if step < WARMUP:
        return LR * (step + 1) / WARMUP                                  # langsam aufwärmen
    progress = (step - WARMUP) / (STEPS - WARMUP)
    return LR * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress)))  # sanft abkühlen


optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=0.1)
start, best = time.time(), float("inf")
for step in range(STEPS + 1):
    for group in optimizer.param_groups:
        group["lr"] = learning_rate(step)
    x, y = get_batch(train)
    loss = loss_on(x, y)
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)              # Notbremse
    optimizer.step()
    if step % 200 == 0:
        val = validation_loss()
        print(f"Schritt {step:5d}  Train {loss.item():.3f}  Valid {val:.3f}  ({time.time() - start:.0f} s)")
        if val < best:                                                   # nur den besten Stand behalten
            best = val
            torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/model.pt")

print(f"Bester Validierungs-Loss {best:.3f}, gespeichert in data/model.pt")

Ausführen

bash
python train.py
text
940,800 Parameter auf cpu
Schritt     0  Train 6.949  Valid 6.942  (1 s)
Schritt   200  Train 4.170  Valid 4.258  (13 s)
Schritt   400  Train 3.875  Valid 3.971  (26 s)
Schritt   600  Train 3.554  Valid 3.771  (39 s)
Schritt   800  Train 3.325  Valid 3.646  (52 s)
Schritt  1000  Train 2.908  Valid 3.546  (65 s)
Schritt  1200  Train 2.845  Valid 3.538  (76 s)
Schritt  1400  Train 2.640  Valid 3.549  (86 s)
Schritt  1600  Train 2.617  Valid 3.535  (95 s)
Schritt  1800  Train 2.565  Valid 3.555  (105 s)
Schritt  2000  Train 2.555  Valid 3.583  (114 s)
Bester Validierungs-Loss 3.535, gespeichert in data/model.pt

Die Ausgabe lesen

Das ist der Moment, in dem Kapitel 23 praktisch wird:

  • Start bei 6,94: genau ln(1024). Das Modell beginnt neutral, so wie check.py es versprochen hat.
  • Schneller Abstieg bis etwa Schritt 1000: Das Modell lernt häufige Wörter, dann typische Wortfolgen, dann Satzbau.
  • Validierung bei 3,54: deutlich unter der Bigram-Baseline von 3,74. Das GPT nutzt also wirklich den längeren Kontext – genau dafür haben wir Attention gebaut.
  • Ab Schritt 1200 stagniert die Validierung, während der Trainings-Loss weiter sinkt. Die Lücke wächst: Das Modell fängt an, die Trainingstexte auswendig zu lernen. Deshalb speichern wir nur den besten Stand.

Warum Dropout? Ein Experiment

Unser Märchenbuch ist klein: 140.000 Tokens für fast eine Million Parameter. Da ist Auswendiglernen verlockend. So sah der gleiche Lauf ohne Dropout aus (Config(..., dropout=0.0)):

Schritt Train Valid
600 3,04 3,66
1000 2,15 3,90
1400 1,15 4,45
2000 0,70 5,00

Der Trainings-Loss stürzt auf 0,7 – das Modell kennt die Trainingstexte fast auswendig. Aber der Validierungs-Loss steigt ab Schritt 600 kräftig an. Klassische Überanpassung aus Kapitel 23, wie aus dem Lehrbuch. Mit Dropout bleibt die Lücke klein, und das beste Modell ist besser: 3,54 statt 3,66.

Probier es selbst: Setz dropout=0.0 und vergleiche. Solche Experimente – eine Sache ändern, Ergebnis messen – sind genau das, was Leute in KI-Laboren den ganzen Tag machen.

Wenn etwas schiefgeht

  • Viel zu langsam (über 10 Sekunden pro 200 Schritte) – Setz STEPS auf 1000 und BATCH auf 16. Das Ergebnis wird etwas schlechter, aber das Prinzip bleibt.
  • Loss wird nan – Die Lernrate ist zu hoch. Probier LR = 1e-3.
  • Valid steigt von Anfang an – Prüfe in get_batch, ob y wirklich um eins verschoben ist.
  • RuntimeError: CUDA out of memory – Auf der Grafikkarte ist zu wenig Platz. BATCH halbieren.

Kurz gemerkt

  • train.py verbindet alles aus Kapitel 22: Batches, AdamW, Warmup und Abkühlen, Clipping, Validierung, Speichern des besten Stands.
  • Der Validierungs-Loss von etwa 3,5 schlägt die Bigram-Baseline von 3,7: Das GPT nutzt den Kontext.
  • Bei wenig Daten lernt ein Modell schnell auswendig. Dropout und „besten Stand speichern“ helfen dagegen.