Tokenwerk · The LLM textbook

Chapter 39 · VII · Build your LLM · 9 minutes

Step 7: From storyteller to chat

With supervised fine-tuning your model learns the conversation format: a request goes in, a fitting answer comes out, then it stops.

What we build

Your model continues text. A chat assistant does something else: it gets a request, responds to it, and then stops. You met exactly this step in chapter 28: supervised fine-tuning, SFT for short.

We build two files:

  • sft.py turns stories into small conversations and keeps training the model on them.
  • chat.py lets you talk to the result.

Honestly, up front: a model with a million parameters won't become all-knowing. But it will learn the format of a conversation: a request comes in, it answers in a fitting style, and it ends the answer by itself. That is the core of SFT.

Where do we get conversations?

For SFT you need pairs of a request and a good answer. Big companies have people write them. We build them from our stories:

  • Request: “Tell me a story that starts like this: Once upon a time, there”, i.e. the first five words of a story.
  • Answer: the rest of that story, here “was a little girl named …”.

That turns our stories into conversations without writing a single one by hand. We take 5,000 of them, which is plenty. Important: we only use stories from the training data (train_docs.json), never from validation.

A conversation in tokens

python
# From text continuer to answerer: supervised fine-tuning (chapter 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(story):
    """A story becomes a conversation: a request with a beginning, the answer is the continuation."""
    words = story.split()
    start = " ".join(words[:5])
    question = f"Tell me a story that starts like this: {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
    # target = input shifted by one; only the answer counts, everything else gets -100
    targets = [-100] * (len(ids) - 1) + answer_ids
    return inputs[: cfg.context], targets[: cfg.context]

The conversation gets the chat template from chapter 28: <bos> <user> request <eos> <assistant> answer <eos>. We shorten the answer so everything fits into 128 tokens.

The crucial part are the targets. There are none for the request: −100 everywhere, meaning “ignore”. Learning only starts at <assistant>. This is what it looks like at the boundary:

Position Input Target
20 “ there” – (ignored)
21 <eos> – (ignored)
22 <assistant> “ was”
23 “ was” “ a”
24 “ a” “ little”

So the model reads the request but is only rated on the answer. The output at <assistant> already has to predict the first answer word, exactly what happens later in the chat.

Training further

python
random.seed(1)
stories = json.loads(open("data/train_docs.json", encoding="utf-8").read())   # training stories only!
examples = [make_example(p) for p in stories[:5000]]
random.shuffle(examples)
print(f"{len(examples)} conversations")


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


optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)       # smaller learning rate: fine-tune, don't relearn
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 step {step:3d}  loss {loss.item():.3f}")

torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/chat.pt")
print("Saved: data/chat.pt")
  • batch pads shorter conversations with <pad>. Their targets are −100, so they don't count.
  • ignore_index=-100 tells cross_entropy: skip these spots.
  • At 0.0003, the learning rate is much smaller than in pretraining. We want to fine-tune, not overwrite everything learned.
  • 300 steps are enough, which takes about a minute. It is saved as data/chat.pt; your story model stays unchanged.

The whole file

python
# From text continuer to answerer: supervised fine-tuning (chapter 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(story):
    """A story becomes a conversation: a request with a beginning, the answer is the continuation."""
    words = story.split()
    start = " ".join(words[:5])
    question = f"Tell me a story that starts like this: {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
    # target = input shifted by one; only the answer counts, everything else gets -100
    targets = [-100] * (len(ids) - 1) + answer_ids
    return inputs[: cfg.context], targets[: cfg.context]


random.seed(1)
stories = json.loads(open("data/train_docs.json", encoding="utf-8").read())   # training stories only!
examples = [make_example(p) for p in stories[:5000]]
random.shuffle(examples)
print(f"{len(examples)} conversations")


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


optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)       # smaller learning rate: fine-tune, don't relearn
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 step {step:3d}  loss {loss.item():.3f}")

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

Chatting

chat.py builds the template exactly like training, just without an answer. The prompt ends with <assistant>, and the model writes until it draws <eos>:

python
# Chatting with your own model (chapter 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("You: ")
    if not question:
        break
    print("Model:", answer(question))

input() waits for your input. An empty line ends the conversation.

Run it

bash
python sft.py
python chat.py
text
...
SFT step 300  loss 2.575
Saved: data/chat.pt
text
You: Tell me a story that starts like this: Tom had a red ball
Model: that he liked to play with his toys. He was always happy and loved to
play with his friends. One day, Tom's mommy said, "Let's play with the dog,
Mia, Tom and Sam are playing ..."

Look closely: the model picks up the beginning (“Tom had a red ball” → “that he liked to play with his toys”), continues in story style and even brings in Tom's mommy and friends. It has learned the format.

And if you ask something that never appeared in the data, like “What is your name?”? Then it answers with story fragments about a boy called Ben. It never learned to talk about itself. That isn't a bug in your code but the most honest lesson from chapter 28: SFT teaches format, not knowledge. What isn't in the data, the model can't do.

If something goes wrong

  • FileNotFoundError: data/train_docs.json: run prepare.py again; older versions didn't write the file.
  • SFT loss becomes nan: a conversation has no active targets left because the request was too long. Check in make_example that the answer isn't empty after shortening.
  • The model never stops: does the last answer token get <eos> as its target? Without + [S["<eos>"]] it doesn't learn to stop.

Key points

  • SFT keeps training the finished model on conversations: <user> request, <assistant> answer.
  • Only the answer is rated. The request gets −100 and is just read.
  • Afterwards the model understands the conversation format and stops by itself. SFT brings no new knowledge.