Tokenwerk · The LLM textbook

Chapter 37 · VII · Build your LLM · 9 minutes

Step 5: Training

The training loop with everything that belongs to it: AdamW, warmup, clipping, validation. After two minutes you have a trained language model.

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

python
# 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

python
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

python
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

python
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

Warmup 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

python
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:

  • AdamW adjusts the step size for every weight individually. weight_decay=0.1 gently 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

python
# 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

bash
python train.py
text
940,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.pt

Reading the output

This is the moment chapter 23 becomes practical:

  • Start at 6.94: exactly ln(1024). The model starts neutral, just as check.py promised.
  • 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 STEPS to 1000 and BATCH to 16. The result gets a bit worse, but the principle stays.
  • Loss becomes nan: the learning rate is too high. Try LR = 1e-3.
  • Valid rises from the start: check in get_batch that y is really shifted by one.
  • RuntimeError: CUDA out of memory: not enough room on the graphics card. Halve BATCH.

Key points

  • train.py connects 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.