What we build
The model is ready; now it should learn. The file train.py is the training loop from chapter 22: draw batches, predict, measure the error, compute backwards, adjust the weights, 2000 times. On top come the pro ingredients from the same chapter: AdamW, warming up and cooling down the learning rate, gradient clipping, regular validation and saving the best state.
On an ordinary laptop this takes one to three minutes. Afterwards you have a trained language model in data/model.pt.
Preparation
# The training loop (chapter 22)
import math
import time
import numpy as np
import torch
import torch.nn.functional as F
from model import GPT, Config
from tokenizer import Tokenizer
STEPS, BATCH, LR, WARMUP = 2000, 32, 2e-3, 100
torch.manual_seed(42)
device = "cuda" if torch.cuda.is_available() else "cpu"
tok = Tokenizer.load("data/tokenizer.json")
train = np.load("data/train.npy")
valid = np.load("data/valid.npy")
cfg = Config(vocab_size=tok.vocab_size)
model = GPT(cfg).to(device)
print(f"{sum(p.numel() for p in model.parameters()):,} parameters on {device}")- The four big knobs are at the very top: 2000 steps, 32 texts per batch, learning rate 0.002, 100 warmup steps.
torch.manual_seed(42)makes the randomness repeatable: same seed, same run.- If you have an NVIDIA graphics card, the script automatically uses
cuda. Otherwise the CPU, which is plenty.
Drawing batches
def get_batch(data):
starts = np.random.randint(0, len(data) - cfg.context - 1, BATCH)
x = np.stack([data[s : s + cfg.context] for s in starts]) # input
y = np.stack([data[s + 1 : s + 1 + cfg.context] for s in starts]) # target: shifted by one
return torch.from_numpy(x).long().to(device), torch.from_numpy(y).long().to(device)Just like the bigram model, only with windows of length cfg.context = 128: random starting points, 128 tokens from each as input and the same 128 tokens shifted by one as the target. So every window contains 128 practice tasks at once.
Measuring the error, honestly too
def loss_on(x, y):
logits = model(x)
return F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1))
@torch.no_grad()
def validation_loss():
model.eval()
losses = [loss_on(*get_batch(valid)).item() for _ in range(20)]
model.train()
return sum(losses) / len(losses)loss_on is the error for one batch. validation_loss measures the error on the validation data, which the model never gets to learn from (chapters 21 and 23):
@torch.no_grad()tells PyTorch: no learning here, so don't remember any calculation path. That saves time and memory.model.eval()switches dropout off,model.train()switches it back on afterwards.- We average over 20 random batches. A single one would be too random.
Planning the learning rate
def learning_rate(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP # warm up slowly
progress = (step - WARMUP) / (STEPS - WARMUP)
return LR * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress))) # cool down gentlyWarmup and cooldown from chapter 22: during the first 100 steps the learning rate slowly rises to 0.002. After that it falls along a cosine curve to a tenth. Big steps at the start, careful ones at the end.
The loop
optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=0.1)
start, best = time.time(), float("inf")
for step in range(STEPS + 1):
for group in optimizer.param_groups:
group["lr"] = learning_rate(step)
x, y = get_batch(train)
loss = loss_on(x, y)
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # emergency brake
optimizer.step()
if step % 200 == 0:
val = validation_loss()
print(f"Step {step:5d} train {loss.item():.3f} valid {val:.3f} ({time.time() - start:.0f} s)")
if val < best: # only keep the best state
best = val
torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/model.pt")
print(f"Best validation loss {best:.3f}, saved to data/model.pt")The heart of it all, and you will recognize every line:
AdamWadjusts the step size for every weight individually.weight_decay=0.1gently pulls all weights toward zero.- Before every step we set the planned learning rate.
- The five core lines: batch, loss,
zero_grad,backward,step. clip_grad_norm_is the emergency brake against huge slopes.- Every 200 steps we measure the validation loss. If it is better than ever before, we save the model. That way the best state is automatically what remains at the end, not the last one.
What gets saved are the settings (cfg.__dict__) and all the weights (state_dict()). Nothing more is needed to rebuild the model later.
The whole file
# The training loop (chapter 22)
import math
import time
import numpy as np
import torch
import torch.nn.functional as F
from model import GPT, Config
from tokenizer import Tokenizer
STEPS, BATCH, LR, WARMUP = 2000, 32, 2e-3, 100
torch.manual_seed(42)
device = "cuda" if torch.cuda.is_available() else "cpu"
tok = Tokenizer.load("data/tokenizer.json")
train = np.load("data/train.npy")
valid = np.load("data/valid.npy")
cfg = Config(vocab_size=tok.vocab_size)
model = GPT(cfg).to(device)
print(f"{sum(p.numel() for p in model.parameters()):,} parameters on {device}")
def get_batch(data):
starts = np.random.randint(0, len(data) - cfg.context - 1, BATCH)
x = np.stack([data[s : s + cfg.context] for s in starts]) # input
y = np.stack([data[s + 1 : s + 1 + cfg.context] for s in starts]) # target: shifted by one
return torch.from_numpy(x).long().to(device), torch.from_numpy(y).long().to(device)
def loss_on(x, y):
logits = model(x)
return F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1))
@torch.no_grad()
def validation_loss():
model.eval()
losses = [loss_on(*get_batch(valid)).item() for _ in range(20)]
model.train()
return sum(losses) / len(losses)
def learning_rate(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP # warm up slowly
progress = (step - WARMUP) / (STEPS - WARMUP)
return LR * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress))) # cool down gently
optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=0.1)
start, best = time.time(), float("inf")
for step in range(STEPS + 1):
for group in optimizer.param_groups:
group["lr"] = learning_rate(step)
x, y = get_batch(train)
loss = loss_on(x, y)
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # emergency brake
optimizer.step()
if step % 200 == 0:
val = validation_loss()
print(f"Step {step:5d} train {loss.item():.3f} valid {val:.3f} ({time.time() - start:.0f} s)")
if val < best: # only keep the best state
best = val
torch.save({"config": cfg.__dict__, "model": model.state_dict()}, "data/model.pt")
print(f"Best validation loss {best:.3f}, saved to data/model.pt")Run it
python train.py940,800 parameters on cpu
Step 0 train 6.953 valid 6.943 (0 s)
Step 200 train 3.764 valid 3.810 (10 s)
Step 400 train 3.395 valid 3.423 (20 s)
Step 600 train 3.248 valid 3.140 (30 s)
Step 800 train 3.145 valid 2.997 (39 s)
Step 1000 train 2.908 valid 2.853 (49 s)
Step 1200 train 2.763 valid 2.762 (59 s)
Step 1400 train 2.765 valid 2.670 (69 s)
Step 1600 train 2.703 valid 2.613 (79 s)
Step 1800 train 2.760 valid 2.548 (89 s)
Step 2000 train 2.713 valid 2.581 (99 s)
Best validation loss 2.548, saved to data/model.ptReading the output
This is the moment chapter 23 becomes practical:
- Start at 6.94: exactly ln(1024). The model starts neutral, just as
check.pypromised. - A fast descent: the model learns frequent words first, then typical word sequences, then sentence structure.
- Validation at 2.55: far below the bigram baseline of 3.7. So the GPT really uses the longer context; that is exactly what we built attention for.
- Training and validation stay close. With 5.5 million tokens there is a lot of data for our million parameters, and in 2000 steps the model sees each token only about one and a half times. The validation loss is even slightly lower than the training loss here, because dropout is only active during training.
Why dropout? An experiment
With lots of data like TinyStories, memorizing is no big danger. That changes with small data. In the German edition of this book the model trains on a single book of fairy tales, only 140,000 tokens. There, the same run without dropout (Config(..., dropout=0.0)) looked like this:
| Step | Train | Valid |
|---|---|---|
| 600 | 3.04 | 3.66 |
| 1000 | 2.15 | 3.90 |
| 1400 | 1.15 | 4.45 |
| 2000 | 0.70 | 5.00 |
The training loss crashes to 0.7: the model knows the training texts almost by heart. But from step 600 the validation loss rises sharply. Textbook overfitting from chapter 23. With dropout the gap stayed small, and the best model was better: 3.54 instead of 3.66.
Try it yourself: train on a small text of your own and set dropout=0.0. Experiments like this (change one thing, measure the result) are exactly what people in AI labs do all day.
If something goes wrong
- Much too slow (over 10 seconds per 200 steps): set
STEPSto 1000 andBATCHto 16. The result gets a bit worse, but the principle stays. - Loss becomes
nan: the learning rate is too high. TryLR = 1e-3. - Valid rises from the start: check in
get_batchthatyis really shifted by one. RuntimeError: CUDA out of memory: not enough room on the graphics card. HalveBATCH.
Key points
train.pyconnects everything from chapter 22: batches, AdamW, warmup and cooldown, clipping, validation, saving the best state.- A validation loss of about 2.5 beats the bigram baseline of 3.7: the GPT uses the context.
- With little data a model quickly memorizes. Dropout and “save the best state” help against that.