Tokenwerk · The LLM textbook

Chapter 38 · VII · Build your LLM · 5 minutes

Step 6: Generating text

Your model writes! With temperature and top-k you produce stories token by token, and see why greedy ends up in loops.

What we build

Your model is trained. Now it should write. The file generate.py loads the saved model and produces text token by token, with temperature and top-k from chapter 24. A little over 30 lines, and at the end your own model writes stories.

Loading the model

python
# Generating text: token by token (chapter 24)
import sys

import torch

from model import GPT, Config
from tokenizer import Tokenizer

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

We rebuild the model with the same settings train.py saved and load the weights into it. model.eval() switches dropout off; while writing we don't want extra randomness inside the model itself.

Writing token by token

python
@torch.no_grad()

This is the loop from chapter 24:

  1. The text starts with <bos> and the prompt, as token IDs.
  2. The model sees at most the last 128 tokens: its context window. Whatever came before is forgotten for this step.
  3. From the output we only take the last position: [0, -1]. It predicts the next token.
  4. Divided by the temperature, then top-k: only the 40 best candidates go into the lottery drum.
  5. torch.multinomial rolls the dice according to the probabilities.
  6. If the drawn token is <eos>, the model itself said “done”; it learned in training where stories end. Otherwise append it and continue.

The last lines read the prompt from the command line. Without a prompt the model starts with “Once upon a time”.

The whole file

python
# Generating text: token by token (chapter 24)
import sys

import torch

from model import GPT, Config
from tokenizer import Tokenizer

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


@torch.no_grad()
def generate(prompt, max_tokens=150, temperature=0.8, top_k=40):
    ids = [tok.special["<bos>"]] + tok.encode(prompt)
    for _ in range(max_tokens):
        x = torch.tensor([ids[-model.cfg.context :]])          # only the context window
        logits = model(x)[0, -1] / temperature                  # last position, temperature
        top = torch.topk(logits, top_k)
        probs = torch.softmax(top.values, dim=-1)               # top-k: only roll among the best
        next_id = top.indices[torch.multinomial(probs, 1)].item()
        if next_id == tok.special["<eos>"]:
            break                                               # the model said “done”
        ids.append(next_id)
    return tok.decode(ids)


torch.manual_seed(0)
prompt = sys.argv[1] if len(sys.argv) > 1 else "Once upon a time"
print(generate(prompt))

Run it

bash
python generate.py "Once upon a time"
text
Once upon a time, there was a little boy named Timmy. Timmy liked to play
with his toys and play with him. One day, Timmy's mommy missed his mum and
daddy said, "Let's go play with Max!" Timmy said Timmy, "Wow, Timmy. Can we
do you mine. That's a very special game to be kind".

An honest look at it:

  • What the model can do: sentence structure, commas, dialogue in quotation marks, typical story words (“Once upon a time”, “One day”), names. Almost all words are real.
  • What it can't do: keep the meaning going across several sentences. Timmy talks to himself, someone suddenly calls him Lily, things don't quite fit. With a million parameters and a few megabytes of text, that is normal. ChatGPT has a hundred thousand times more parameters and has read millions of books.

Try your own beginnings: “The little dog”, “Lily wanted to”, “One day, a big bear”. Every start gives a different story.

Experiments with temperature

Change the temperature in generate and compare. With the same beginning and the same seed, I got:

Setting Result
temperature=0.3 Once upon a time, there was a little girl named Lily. She loved to play with her toys and play with her friends. …
temperature=0.8 Once upon a time, there was a little boy named Timmy. Timmy liked to play with his toys and play with him. …
temperature=1.5 Once upon a time, there was a very small bird named Lily. Max was three years old and liked vegetables who liked dressed animals around everything. …
top_k=1 (greedy) Once upon a time, there was a little girl named Lily. She loved to play with her toys and play with her toys. …

Just as chapter 24 described: low temperature gets cautious and safe. High temperature gets wild: a bird named Lily who is suddenly Max and three years old. And greedy, always taking the most likely token, starts repeating itself right away (“play with her toys and play with her toys”). 0.8 is a good compromise.

If something goes wrong

  • FileNotFoundError: data/model.pt: run python train.py first.
  • RuntimeError: Error(s) in loading state_dict: model.py was changed after training. Either undo the change or train again.
  • The text ends immediately: the model drew <eos> early. Just run it again or give a longer beginning.
  • Only letter salad: training went wrong. Is the best validation loss above 4? Then go back to step 5.

Key points

  • Generating means: predict, roll the dice, append, repeat, until <eos> or the maximum length.
  • Only the last position counts, and the model sees at most its context window.
  • Temperature and top-k decide between boring, usable and wild. Greedy quickly ends up repeating itself.