7. Train and evaluate your GPT
In this chapter
- The language-modeling objective: cross-entropy on the next token, and what its numbers mean.
- A complete training loop with AdamW, warmup, cosine decay, gradient clipping and periodic evaluation.
- Sanity checks that catch most training bugs in minutes.
- Overfitting, seen live on a real text, and what it does and doesn't tell you.
You will build
engine/train.py: lm_loss, evaluate and train. You'll train your Chapter 6 GPT on The Verdict and sample from it.
Time: 4-6 hours. GPU: helpful but not needed (the default experiment takes under a minute on a laptop CPU).
The objective: predict the next token
Chapter 4 turned text into windows of inputs and targets shifted by one. For every position $t$ in every window, the model outputs logits $z$, a score per vocabulary entry, and we want the probability of the true next token $y$ to be high. Softmax turns logits into probabilities, and the cross-entropy loss is the negative log of the probability assigned to the correct answer:
$$ p(y) = \frac{e^{z_y}}{\sum_j e^{z_j}}, \qquad L = -\log p(y). $$
The numbers are worth internalizing:
| probability of the right token | loss |
|---|---|
| 1.00 | 0.000 |
| 0.50 | 0.693 |
| 0.10 | 2.303 |
| 0.01 | 4.605 |
| 1/V (uniform guessing, V = 512) | 6.238 |
The last row is the most useful number in a training log. An untrained model should start at about $\log V$. If your first loss is 30, the initialization or the loss computation is broken. If it’s 0.5, the model can see its targets.
The gradient has a beautiful form
The derivative of cross-entropy with respect to each logit is
$$ \frac{\partial L}{\partial z_j} = p_j - \mathbb 1[j = y]. $$
If the target has probability 0.2, its logit’s gradient is −0.8, so gradient descent pushes it up. A wrong token with probability 0.3 gets +0.3 and is pushed down. Every logit is nudged in proportion to how wrong its probability is. This gradient then flows back through the head, every block, and the embeddings, via the chain rule from Chapter 3.
Use logits, not probabilities
PyTorch’s F.cross_entropy takes logits and computes a numerically stable log-softmax internally. Passing it probabilities (applying softmax yourself first) silently computes a different, wrong loss. It also wants 2-D logits and 1-D targets, so flatten batch and time together:
def lm_loss(model, x, y):
"""Mean next-token cross-entropy over every position. (Your engine: Chapter 7)
logits [B, T, V] and targets [B, T] are flattened to [B*T, V] and [B*T]. cross_entropy
applies a stable log-softmax itself, so it takes raw logits, never probabilities.
"""
logits = model(x)
return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), y.reshape(-1))
For a validation set, average by tokens, not by batches. If batches have different sizes, a plain average of batch means over-weights the small ones:
@torch.no_grad()
def evaluate(model, x, y, batch_size=32):
"""Token-weighted mean loss over a whole split, in eval mode, without gradients. (Your engine: Chapter 7)"""
was_training = model.training
model.eval()
total, count = 0.0, 0
for start in range(0, len(x), batch_size):
xb, yb = x[start:start + batch_size], y[start:start + batch_size]
total += lm_loss(model, xb, yb).item() * yb.numel()
count += yb.numel()
model.train(was_training)
return total / count
The training loop
Here’s the full loop you’ll write. It’s the five-line core from Chapter 3, plus the details that make training stable:
def train(model, train_xy, valid_xy, steps, batch_size=16, lr=3e-3, weight_decay=0.1,
warmup=20, clip=1.0, eval_every=50, seed=0, log=print):
"""AdamW with linear warmup then cosine decay, gradient clipping and periodic evaluation. (Your engine: Chapter 7)
Returns a list of {"step", "train_loss", "valid_loss", "lr"} records.
"""
device = next(model.parameters()).device
tx, ty = (t.to(device) for t in train_xy)
vx, vy = (t.to(device) for t in valid_xy)
generator = torch.Generator().manual_seed(seed)
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay, betas=(0.9, 0.95))
def schedule(step):
if step < warmup:
return (step + 1) / warmup
progress = (step - warmup) / max(1, steps - warmup)
return 0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress))
history = []
for step in range(steps):
for group in optimizer.param_groups:
group["lr"] = lr * schedule(step)
model.train()
index = torch.randint(len(tx), (batch_size,), generator=generator).to(device)
loss = lm_loss(model, tx[index], ty[index])
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
optimizer.step()
if step % eval_every == 0 or step + 1 == steps:
record = {"step": step + 1, "train_loss": round(loss.item(), 4),
"valid_loss": round(evaluate(model, vx, vy), 4),
"lr": round(optimizer.param_groups[0]["lr"], 6)}
history.append(record)
if log:
log(json.dumps(record))
return history
- AdamW with $\beta_2 = 0.95$ (rather than the default 0.999) is a common choice for transformers. It adapts faster when gradient statistics change.
- Warmup ramps the learning rate up over the first steps. Adam’s early estimates of gradient variance are noisy, so the first updates should be small.
- Cosine decay lowers the learning rate smoothly to 10% of its peak, so the model settles into a minimum instead of bouncing around it.
- Gradient clipping caps the global gradient norm at 1.0, so one unusual batch can’t throw the weights far off course.
- Evaluation runs on held-out data, in
eval()mode, without gradients, everyeval_everysteps.
First, prove the loop can learn
Before training on real data, run two checks that catch most bugs in minutes:
- Initial loss ≈ log V. It confirms the model and loss are wired correctly.
- Overfit one batch. Train repeatedly on a single fixed batch. A correct model and loop drive the loss near zero within a couple of hundred steps. If they don’t, there’s a bug: targets not shifted (or shifted twice), the optimizer created before parameters were replaced,
zero_gradmissing, the head disconnected, or a learning rate far too large.
Both are in your milestone tests. A model that can’t memorize one batch isn’t ready for a dataset.
Experiment 1: a task the model can master
The synthetic colors corpus repeats two phrases, “red green blue” and “blue green red”, 160 times each. With a character tokenizer and a 2-layer model, validation loss falls right along with training loss (python run.py train --corpus colors --steps 200 --context 32 --width 64 --layers 2):
{"step": 121, "train_loss": 0.0928, "valid_loss": 0.0676, ...}
{"step": 161, "train_loss": 0.0688, "valid_loss": 0.0636, ...}
{"step": 200, "train_loss": 0.0687, "valid_loss": 0.053, ...}
The validation text is new, but its patterns are the same as in training, so what the model learned transfers. This is what successful generalization looks like, on a deliberately easy task.
Experiment 2: a real text, and overfitting
Now train on The Verdict with a 512-token BPE vocabulary, 64-token windows and a 4-layer, 128-wide model of 867,072 parameters (python run.py train --steps 600). The 90/10 split gives 260 training windows and 28 validation windows. On a laptop CPU this takes about 40 seconds:
| step | train loss | valid loss |
|---|---|---|
| 1 | 6.248 | 6.201 |
| 61 | 4.833 | 5.016 |
| 121 | 3.849 | 4.424 |
| 181 | 3.100 | 4.434 |
| 241 | 2.628 | 4.704 |
| 361 | 1.599 | 5.534 |
| 481 | 0.717 | 6.099 |
| 600 | 0.437 | 6.410 |
The run starts at the uniform loss, 6.24 = log 512, as it should. Validation loss is best around step 120-180 and then gets worse, while training loss keeps falling toward zero. This is overfitting. The story is about 7,000 tokens long, and the model has 867k parameters, more than a hundred per training token. After learning the general statistics of English, its cheapest way to keep lowering the training loss is to memorize the training text verbatim. That memorized text is useless, even harmful, on the unseen 10%.
Samples make this concrete. At step 600, prompted with the story’s first words, the model reproduces the opening sentence and then degenerates into fragments of memorized phrases:
I HAD always thought Jack Gisburn rather a cheap genies be sply oweagal note that Emperors of thereerly a
Stopping at step 150, near the validation minimum, gives text that is less memorized but still clearly not fluent English. 7,000 tokens is far too little to learn a language from:
I HAD always thought, and he was the coree to doree of the coree to dorethethetheting, and the coreting
BALLM Chapter 5 shows the same effect on the same text. There’s nothing wrong with the code. The remedy is data, not cleverness: real pretraining uses trillions of tokens, so every token is seen once or a handful of times and memorization is rare. Regularization (dropout, used here at 0.1, and weight decay) and early stopping (keeping the checkpoint with the best validation loss) help at the margin.
Note
Perplexity is $e^{\text{loss}}$: here 512 at the start (uniform over 512 tokens), 83 at the validation minimum, and 1.5 on the memorized training set. Read it as “the model is as uncertain as if it were choosing uniformly among this many tokens”. Perplexities are only comparable between models with the same tokenizer: a 512-token vocabulary and a 150,000-token one measure uncertainty per very different units.
Teacher forcing and free-running generation
During training, every position sees the true previous tokens, even where the model would have predicted something else. This is teacher forcing, and it’s why one forward pass gives $T$ training examples. During generation, the model sees its own previous outputs. One early mistake changes every later input, and the model was never trained on contexts containing its own errors. So low validation loss is necessary but not sufficient for good generations. Always look at samples as well as loss curves.
Saving and loading
A checkpoint for inference needs the weights and the configuration that defines the architecture, plus whatever reconstructs the tokenizer:
def save_checkpoint(path, model, extra=None):
"""Weights + config + anything needed to rebuild the tokenizer. Not an exact-resume record:
optimizer and RNG state are deliberately omitted."""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
torch.save({"config": vars(model.cfg), "state": model.state_dict(), "extra": extra or {}}, path)
def load_checkpoint(path, model_class, config_class, device="cpu"):
"""Only load checkpoints you created: torch.load of untrusted pickles can run code."""
blob = torch.load(path, map_location=device, weights_only=True)
model = model_class(config_class(**blob["config"])).to(device)
model.load_state_dict(blob["state"])
return model.eval(), blob["extra"]
Exactly resuming training also needs the optimizer state (Adam’s $m$ and $v$), the learning-rate schedule position, and the random-number-generator states. Our checkpoints deliberately omit those.
Warning
torch.loadon.pt/.binfiles unpickles Python objects, and a malicious pickle can run arbitrary code. Load only files you created, withweights_only=True. Real model distribution has moved to safetensors (Chapter 9), which is plain bytes plus a JSON header.
How big is real pretraining?
A useful rule of thumb: training costs about $6ND$ floating-point operations for $N$ parameters and $D$ training tokens. That’s 2 per parameter per token for the forward pass, and 4 for the backward. For our toy model, $6 \times 867\text{k} \times 600 \text{ steps} \times 16 \times 64$ tokens is about 3.2 TFLOP, a few seconds of a laptop. Qwen3-0.6B was trained on around 36 trillion tokens: about $1.3 \times 10^{23}$ FLOPs, or tens of thousands of GPU-days. This is why you’ll load pretrained weights from Chapter 9 on, and why fine-tuning (Part V) changes a trained model rather than starting over.
Build it
Engine milestone 7: train a GPT. In engine/train.py, implement lm_loss, evaluate and train (save/load helpers are provided).
pytest tests/test_ch07_training.py
python run.py train --impl engine --steps 600
python run.py generate --impl engine
The tests check the loss against a hand-computed cross-entropy, check that the initial loss is near log V, overfit one batch, and train on the colors task until validation loss halves.
Stretch exercises
- ★ Add early stopping to
train: keep a copy of the weights with the lowest validation loss and restore them at the end. How much better do the samples look? Where:traininengine/train.py. - ★★ Change one variable at a time (dropout 0.0 vs 0.3, width 64 vs 256, learning rate 1e-3 vs 1e-2) and record the best validation loss of each run. Predict each effect before you run it. Where:
experiments/ch07.py(create it), adaptingrun.py’scmd_trainand itsGPTConfig. - ★★ Train on a larger public-domain text (any Project Gutenberg novel, or several) with the same model. How does the gap between train and validation loss change with 10x more data? Where: put the new corpus in
data/; load it inexperiments/ch07.py(create it) usingengine.dataandengine.train. - ★★★ Implement gradient accumulation: compute gradients over $k$ small batches before each optimizer step, weighting by tokens. Verify that $k$ micro-batches of size 8 give the same update as one batch of 8k. Where:
traininengine/train.py.
Check your understanding
- Why does cross-entropy take logits rather than probabilities?
- What initial loss do you expect for a vocabulary of 50,257, and what does a much larger value suggest?
- Why can validation loss rise while training loss falls?
- Why can a model with low teacher-forced loss still produce poor generations?
- What extra state, beyond the weights, do you need to resume training exactly?
Going deeper
- BALLM Chapter 5 §§5.1-5.2 (pp. 128-150): evaluating generative models, training on The Verdict and observing overfitting. Appendix D (pp. 313-321): warmup, cosine decay, clipping.
- Kaplan et al., Scaling Laws for Neural Language Models (2020) and Hoffmann et al., Training Compute-Optimal Large Language Models (“Chinchilla”, 2022): how parameters, data and compute trade off.
- Andrej Karpathy, Let’s reproduce GPT-2 (124M) (video): a full pretraining run of the model your Chapter 9 loader will import.