Tokenwerk · Das LLM-Lehrbuch

Kapitel 39 · VII · Bau dein LLM · 9 Minuten

Schritt 7: Vom Märchenerzähler zum Chat

Mit Supervised Fine-Tuning lernt dein Modell das Gesprächsformat: Bitte rein, passende Antwort raus, dann Schluss.

Was wir bauen

Dein Modell setzt Text fort. Ein Chat-Assistent macht aber etwas anderes: Er bekommt eine Bitte und antwortet darauf – und hört dann auf. Genau diesen Schritt hast du in Kapitel 28 kennengelernt: Supervised Fine-Tuning, kurz SFT.

Wir bauen zwei Dateien:

  • sft.py macht aus Märchenabsätzen kleine Gespräche und trainiert das Modell darauf weiter.
  • chat.py lässt dich mit dem Ergebnis reden.

Ehrlich vorab: Aus einem Modell mit einer Million Parametern wird kein Allwissender. Es wird aber lernen, das Format eines Gesprächs zu verstehen: Eine Bitte kommt, dann antwortet es im passenden Stil und beendet die Antwort selbst. Genau das ist der Kern von SFT.

Woher nehmen wir Gespräche?

Für SFT braucht man Paare aus Bitte und guter Antwort. Große Firmen lassen sie von Menschen schreiben. Wir basteln sie aus unseren Märchen:

  • Bitte: „Erzähl ein Märchen, das so beginnt: Das Mädchen aber tat wie“ – also die ersten fünf Wörter eines Absatzes.
  • Antwort: der Rest dieses Absatzes, hier „die Haulemännerchen gesagt …“.

So entstehen aus 358 Trainingsabsätzen 358 Gespräche, ohne dass wir eines von Hand schreiben müssen. Wichtig: Wir nehmen nur Absätze aus den Trainingsdaten (train_docs.json), nie aus der Validierung.

Ein Gespräch in Tokens

python
# Vom Textfortsetzer zum Antwortgeber: Supervised Fine-Tuning (Kapitel 28)
import json
import random

import torch
import torch.nn.functional as F

from model import GPT, Config
from tokenizer import Tokenizer

tok = Tokenizer.load("data/tokenizer.json")
S = tok.special
checkpoint = torch.load("data/model.pt")
cfg = Config(**checkpoint["config"])
model = GPT(cfg)
model.load_state_dict(checkpoint["model"])


def make_example(paragraph):
    """Aus einem Märchenabsatz wird ein Gespräch: Frage nach einem Anfang, Antwort = Fortsetzung."""
    words = paragraph.split()
    start = " ".join(words[:5])
    question = f"Erzähl ein Märchen, das so beginnt: {start}"
    answer = " ".join(words[5:])
    ids = [S["<bos>"], S["<user>"]] + tok.encode(question) + [S["<eos>"], S["<assistant>"]]
    answer_ids = tok.encode(" " + answer)[: cfg.context - len(ids)] + [S["<eos>"]]
    inputs = ids + answer_ids
    # Ziel = Eingabe um eins verschoben; nur die Antwort zählt, alles andere bekommt -100
    targets = [-100] * (len(ids) - 1) + answer_ids
    return inputs[: cfg.context], targets[: cfg.context]

Das Gespräch bekommt die Chatvorlage aus Kapitel 28: <bos> <user> Bitte <eos> <assistant> Antwort <eos>. Die Antwort kürzen wir so, dass alles in 128 Tokens passt.

Der entscheidende Teil sind die Ziele. Für die Bitte gibt es keine – überall steht −100, also „ignorieren“. Erst ab <assistant> wird gelernt. So sieht das an der Grenze aus:

Position Eingabe Ziel
26 „ wie“ – (ignoriert)
27 <eos> – (ignoriert)
28 <assistant> „ die“
29 „ die“ „ H“
30 „ H“ „au“

Das Modell liest die Bitte also mit, wird aber nur für die Antwort bewertet. Die Ausgabe auf <assistant> muss schon das erste Antwortwort vorhersagen – genau das, was später im Chat passiert.

Weitertrainieren

python
random.seed(1)
paragraphs = json.loads(open("data/train_docs.json", encoding="utf-8").read())   # nur Trainingsabsätze!
examples = [make_example(p) for p in paragraphs]
random.shuffle(examples)
print(f"{len(examples)} Gespräche, Beispiel-Frage:", tok.decode(examples[0][0][2:20]))


def batch(items):
    n = max(len(x) for x, _ in items)
    x = torch.tensor([xi + [S["<pad>"]] * (n - len(xi)) for xi, _ in items])   # auffüllen ...
    y = torch.tensor([yi + [-100] * (n - len(yi)) for _, yi in items])          # ... und ignorieren
    return x, y


optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)       # kleinere Lernrate: nachjustieren, nicht neu lernen
model.train()
for step in range(301):
    x, y = batch(random.sample(examples, 16))
    logits = model(x)
    loss = F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1), ignore_index=-100)
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()
    if step % 50 == 0:
        print(f"SFT-Schritt {step:3d}  Loss {loss.item():.3f}")

torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/chat.pt")
print("Gespeichert: data/chat.pt")
  • batch füllt kürzere Gespräche mit <pad> auf. Deren Ziele sind −100, sie zählen also nicht.
  • ignore_index=-100 sagt cross_entropy: Diese Stellen überspringen.
  • Die Lernrate ist mit 0,0003 deutlich kleiner als beim Vortraining. Wir wollen nachjustieren, nicht alles Gelernte überschreiben.
  • 300 Schritte reichen, das dauert etwa eine Minute. Gespeichert wird als data/chat.pt – dein Märchenmodell bleibt unverändert.

Die ganze Datei

python
# Vom Textfortsetzer zum Antwortgeber: Supervised Fine-Tuning (Kapitel 28)
import json
import random

import torch
import torch.nn.functional as F

from model import GPT, Config
from tokenizer import Tokenizer

tok = Tokenizer.load("data/tokenizer.json")
S = tok.special
checkpoint = torch.load("data/model.pt")
cfg = Config(**checkpoint["config"])
model = GPT(cfg)
model.load_state_dict(checkpoint["model"])


def make_example(paragraph):
    """Aus einem Märchenabsatz wird ein Gespräch: Frage nach einem Anfang, Antwort = Fortsetzung."""
    words = paragraph.split()
    start = " ".join(words[:5])
    question = f"Erzähl ein Märchen, das so beginnt: {start}"
    answer = " ".join(words[5:])
    ids = [S["<bos>"], S["<user>"]] + tok.encode(question) + [S["<eos>"], S["<assistant>"]]
    answer_ids = tok.encode(" " + answer)[: cfg.context - len(ids)] + [S["<eos>"]]
    inputs = ids + answer_ids
    # Ziel = Eingabe um eins verschoben; nur die Antwort zählt, alles andere bekommt -100
    targets = [-100] * (len(ids) - 1) + answer_ids
    return inputs[: cfg.context], targets[: cfg.context]


random.seed(1)
paragraphs = json.loads(open("data/train_docs.json", encoding="utf-8").read())   # nur Trainingsabsätze!
examples = [make_example(p) for p in paragraphs]
random.shuffle(examples)
print(f"{len(examples)} Gespräche, Beispiel-Frage:", tok.decode(examples[0][0][2:20]))


def batch(items):
    n = max(len(x) for x, _ in items)
    x = torch.tensor([xi + [S["<pad>"]] * (n - len(xi)) for xi, _ in items])   # auffüllen ...
    y = torch.tensor([yi + [-100] * (n - len(yi)) for _, yi in items])          # ... und ignorieren
    return x, y


optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)       # kleinere Lernrate: nachjustieren, nicht neu lernen
model.train()
for step in range(301):
    x, y = batch(random.sample(examples, 16))
    logits = model(x)
    loss = F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1), ignore_index=-100)
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()
    if step % 50 == 0:
        print(f"SFT-Schritt {step:3d}  Loss {loss.item():.3f}")

torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/chat.pt")
print("Gespeichert: data/chat.pt")

Chatten

chat.py baut die Vorlage genauso auf wie das Training – nur ohne Antwort. Der Prompt endet mit <assistant>, und das Modell schreibt los, bis es <eos> zieht:

python
# Mit dem eigenen Modell chatten (Kapitel 28)
import torch

from model import GPT, Config
from tokenizer import Tokenizer

tok = Tokenizer.load("data/tokenizer.json")
S = tok.special
checkpoint = torch.load("data/chat.pt")
model = GPT(Config(**checkpoint["config"]))
model.load_state_dict(checkpoint["model"])
model.eval()


@torch.no_grad()
def answer(question, max_tokens=120, temperature=0.7, top_k=40):
    ids = [S["<bos>"], S["<user>"]] + tok.encode(question) + [S["<eos>"], S["<assistant>"]]
    start = len(ids)
    for _ in range(max_tokens):
        logits = model(torch.tensor([ids[-model.cfg.context :]]))[0, -1] / temperature
        top = torch.topk(logits, top_k)
        next_id = top.indices[torch.multinomial(torch.softmax(top.values, -1), 1)].item()
        if next_id == S["<eos>"]:
            break
        ids.append(next_id)
    return tok.decode(ids[start:])


torch.manual_seed(0)
while True:
    question = input("Du: ")
    if not question:
        break
    print("Modell:", answer(question))

input() wartet auf deine Eingabe. Eine leere Zeile beendet das Gespräch.

Ausführen

bash
python sft.py
python chat.py
text
...
SFT-Schritt 300  Loss 1.996
Gespeichert: data/chat.pt
text
Du: Erzähl ein Märchen, das so beginnt: Es war einmal eine arme Witwe
Modell: , die hatte zwei Tochter. Einmal sind alle Tage, die nicht Kinder,
die wüße ihre Kinder, ... und sprach: »Wenn ich dich feinertig, so will ich
dich selbst in den Wald gehen.«

Schau genau hin: Das Modell greift den Anfang auf („Es war einmal eine arme Witwe“ → „die hatte zwei Tochter“ – fast richtig), schreibt im Märchenstil weiter und bringt sogar eine wörtliche Rede unter. Es hat das Format gelernt.

Und wenn du etwas fragst, das nie in den Daten vorkam, etwa „Wie heißt du?“? Dann antwortet es – mit Märchenfetzen über Rapunzel. Es hat nie gelernt, über sich selbst zu sprechen. Das ist kein Fehler deines Codes, sondern die ehrlichste Lektion aus Kapitel 28: SFT bringt Format bei, kein Wissen. Was nicht in den Daten steckt, kann das Modell nicht.

Wenn etwas schiefgeht

  • FileNotFoundError: data/train_docs.json – prepare.py noch einmal ausführen; ältere Fassungen haben die Datei nicht geschrieben.
  • SFT-Loss wird nan – Ein Gespräch hat gar keine aktiven Ziele mehr, weil die Bitte zu lang war. Prüfe in make_example, dass die Antwort nach dem Kürzen nicht leer ist.
  • Das Modell hört nie auf – Bekommt das letzte Antworttoken <eos> als Ziel? Ohne + [S["<eos>"]] lernt es das Aufhören nicht.

Kurz gemerkt

  • SFT trainiert das fertige Modell mit Gesprächen weiter: <user> Bitte, <assistant> Antwort.
  • Nur die Antwort wird bewertet. Die Bitte bekommt −100 und wird bloß gelesen.
  • Danach versteht das Modell das Gesprächsformat und hört selbst auf. Neues Wissen bringt SFT nicht.