Tokenwerk · Das LLM-Lehrbuch

Kapitel 22 · IV · Trainieren und prüfen · 6 Minuten

Deine erste Trainingsschleife

Batch ziehen, vorhersagen, Fehler ableiten, Gewichte ändern. Dazu eine brauchbare Lernrate, saubere Validierung und sichere Zwischenstände.

Die Lernschleife kennst du bereits

Beim kleinen Netz haben wir eine Vorhersage berechnet, ihren Fehler bestimmt, rückwärts die Ableitungen gefunden und die Parameter etwas verändert. Hier machen wir dieselben vier Schritte für viele Textbeispiele. Neu ist die Organisation: Welche Beispiele kommen wann dran, wie groß ist ein Schritt und wann speichern wir den Zustand?

Ein Optimierer ist die Regel, nach der Ableitungen in Parameteränderungen übersetzt werden. Unser einfacher Gradientenschritt ist ein Optimierer. AdamW führt zusätzlich laufende Statistiken früherer Gradienten pro Parameter und verwendet sie zur Anpassung der Schrittgrößen. Außerdem zieht er bei jedem Schritt einen kleinen Anteil jedes Gewichts ab (Weight Decay), sodass große Gewichte absolut stärker schrumpfen. Das ersetzt nicht das eigentliche Lernen am Vorhersagefehler.

Eine Lernratenplanung (englisch „Schedule“) verändert die Lernrate über die Zeit. Beim Warmup steigt sie anfangs allmählich. Ein anschließender Cosine-Schedule senkt sie entlang einer glatten Kosinuskurve. Du musst diese Kurve nicht auswendig kennen: Prüfe zuerst, ob dein Lernlauf mit einer kleinen festen Lernrate funktioniert.

Ein Batch enthält mehrere Vorhersageaufgaben

Der Trainer zieht BB Startpositionen aus der vorbereiteten Trainingstokenfolge. Von jeder Startposition nimmt er T+1T+1 Tokens. Die ersten TT sind Input, die letzten TT sind Target. So entstehen gleichlange Arrays ohne Padding.

python
starts = rng.integers(0, len(data) - context, size=batch_size)
chunks = np.stack([data[i:i+context+1] for i in starts])
x = torch.from_numpy(chunks[:, :-1].copy()).long().to(device)
y = torch.from_numpy(chunks[:, 1:].copy()).long().to(device)

copy() verhindert Probleme mit bestimmten Arrayansichten. Die Datenmenge muss länger als das Kontextfenster sein. Ein sinnvoller Trainer prüft das früh mit einer verständlichen Fehlermeldung.

Die fünf Kernzeilen

python
optimizer.zero_grad(set_to_none=True)
logits = model(x)
loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1))
loss.backward()
optimizer.step()

Das ist das Herz des Trainings, aber noch kein kompletter Trainingsbetrieb. Du brauchst Fehlermeldungen für NaNs, eine kontrollierte Lernrate, Zwischenstände und Validierung. Diese Ergänzungen dienen der Nachvollziehbarkeit, nicht einer möglichst großen Zahl von Optionen.

Lernrate und Warmup

Eine feste Lernrate kann bei einem kleinen Modell gut funktionieren. Unser Referenztrainer verwendet eine kurze lineare Aufwärmphase und danach einen Cosine-Abfall bis zu einem Mindestanteil. Warmup reduziert die ersten Schritte, während Aktivierungen und Optimizer-Zustände sich einpendeln.

Hier bezeichnet s die Nummer des gerade ausgeführten Schritts, beginnend bei 0. Bei einem geplanten Lauf mit SS Schritten und WW Warmup-Schritten steigt die Rate zunächst proportional zu (s+1)/W(s+1)/W. Danach berechnet ein Fortschrittswert zwischen null und eins den Cosine-Verlauf. Der Plan ist an die gewünschte Gesamtzahl der Schritte gebunden. Wenn du bei einer Fortsetzung die Gesamtzahl änderst, änderst du auch den restlichen Lernratenplan.

Das Browserexperiment zeigt die geplante Rate, keinen gemessenen Loss. Diese Unterscheidung ist wichtig: Eine hübsche Cosine-Kurve beweist keinen Lernfortschritt.

Gradient Accumulation

Wenn ein großer Batch nicht in den Speicher passt, kannst du mehrere kleine Microbatches nacheinander vorwärts und rückwärts rechnen, bevor du den Optimizer aufrufst. Für gleich große und gleich gewichtete Microbatches mit Anzahl AA teilst du den gemittelten Loss jeweils durch AA.

python
optimizer.zero_grad(set_to_none=True)
for _ in range(accum_steps):
    x, y = get_batch()
    loss = F.cross_entropy(model(x).reshape(-1, V), y.reshape(-1))
    (loss / accum_steps).backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()

B zählt Beispiele pro Microbatch, T die Tokens pro Beispiel und A die angesammelten Microbatches. Die effektive Batchgröße an Tokens ist BTABTA. Zum Beispiel ergeben 2 × 4 × 3 insgesamt 24 Tokens pro gemeinsamer Aktualisierung. Die Aktivierungen eines Microbatches können nach seinem Backward freigegeben werden. Die Gradienten sammeln sich in denselben Parameterfeldern.

Gradient Clipping

Clipping begrenzt hier die globale Gradientennorm. Eine plötzliche sehr große Ableitung kann sonst einen destruktiven Schritt verursachen. Clipping ersetzt aber nicht die Fehlersuche: Wenn fast jeder Schritt stark begrenzt wird oder NaNs auftreten, prüfe Lernrate, Daten, Initialisierung und numerische Operationen.

Clipping nach jedem Microbatch wäre nicht dasselbe wie Clipping der aufsummierten Gradienten. Unser Pretraining clippt erst nach dem vollständigen Accumulation-Schritt.

Ein konkreter Lauf

bash
python train.py --data runs/data --out runs/base \
  --steps 2000 --context 128 --dim 128 --layers 4 \
  --heads 4 --batch 8 --accum 2 --lr 0.0003

Diese Parameter sind ein Ausgangspunkt, keine versprochene optimale Einstellung. Bei 8⋅128⋅2=20488\cdot128\cdot2=2048 Tokens pro Optimizer-Schritt ergeben 2000 Schritte ungefähr 4,096 Millionen verarbeitete Tokens. Wenn der Korpus nur 100.000 Tokens besitzt, werden viele Beispiele mehrfach benutzt.

Die Laufzeit hängt erheblich von Backend, Gerät, Datengröße und Implementierung ab. Miss Tokens pro Sekunde auf deiner Maschine. Eine pauschale Minutenzahl aus der Parameterzahl allein wäre unzuverlässig.

Validierung ohne versehentliche Zustandsänderung

Bei der Validierung setzen wir Eval-Modus, deaktivieren Gradienten und berechnen auf getrennten Daten den mittleren Loss. Der Trainer verwendet eine eigene Zufallsquelle für Validierungsfenster, damit eine Evaluation nicht die Reihenfolge der zukünftigen Trainingsbatches verändert.

Einige Fenster liefern nur eine Stichprobe. Die Ausgabe nennt deshalb einen geschätzten Validierungs-Loss. Für eine endgültige Messung bewertest du die gesamte Validierungsmenge mit definierter Kontextbehandlung. Verschiedene Fensterverfahren können unterschiedlich viele Kontexttokens für dieselben Ziele anbieten und sind daher nicht völlig austauschbar.

Checkpoints und Fortsetzen

last.pt speichert das Modell, die Konfiguration, Tokenizer, Optimizer, Schrittzahl und RNG-Zustände. Es wird atomar über eine temporäre Datei ersetzt, damit ein unterbrochener Schreibvorgang nicht den vorherigen Zwischenstand überschreibt. best.pt speichert den besten beobachteten Validierungszwischenstand des Laufs.

bash
python train.py --data runs/data --out runs/base --steps 2000 --resume runs/base/last.pt

Die Fortsetzung prüft die Datenidentität. Auf derselben Umgebung ist ein sehr ähnlicher Verlauf erreichbar; zwischen unterschiedlichen Beschleunigern oder Bibliotheksversionen ist bitgenaue Wiederholung nicht garantiert. Das Projekt ist kein verteiltes Fehlertoleranzsystem, sondern ein lokal nachvollziehbarer Trainer.