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
# 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
@torch.no_grad()This is the loop from chapter 24:
- The text starts with
<bos>and the prompt, as token IDs. - The model sees at most the last 128 tokens: its context window. Whatever came before is forgotten for this step.
- From the output we only take the last position:
[0, -1]. It predicts the next token. - Divided by the temperature, then top-k: only the 40 best candidates go into the lottery drum.
torch.multinomialrolls the dice according to the probabilities.- 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
# 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
python generate.py "Once upon a time"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: runpython train.pyfirst.RuntimeError: Error(s) in loading state_dict:model.pywas 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.