What we build
Your GPT roughly matches GPT-2 from 2019. In chapter 25 you met three components that are in almost every large model today: RMSNorm, RoPE and SwiGLU. Now you build them in, and measure whether they really help with our stories.
So you can compare both variants, Config gets a switch: modern=False is your model so far, modern=True the modern version. Everything else stays the same.
The switch
@dataclassRMSNorm
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weightLike LayerNorm, just without subtracting the average (chapter 25): average the squares, take the square root, divide by that size, multiply by a learnable weight. torch.rsqrt means “1 divided by the square root”, one calculation instead of two. eps is the safety number against dividing by zero.
RoPE: position by rotating
def rope(x):
"""Rotate pairs of numbers depending on position (x: [B, H, T, d])."""
T, d = x.shape[-2], x.shape[-1]
speed = 10000 ** (-torch.arange(0, d, 2, device=x.device) / d) # every pair turns at its own speed
angle = torch.arange(T, device=x.device)[:, None] * speed[None, :] # [T, d/2]
cos, sin = angle.cos(), angle.sin()
even, odd = x[..., 0::2], x[..., 1::2]
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)This is the clock-hand idea from chapter 25 in seven lines:
speedsets how fast each pair of numbers turns: the first pair like a second hand, the last like an hour hand.angleis position times speed: at position 0 nothing turns, at position 10 it turns ten times as far as at position 1.evenandoddare the two numbers of each pair. The line withcosandsinis the rotation formula from the math corner of chapter 25.stackandflattenput the rotated pairs back in their original order.
Only query and key are rotated, not value. That is why attention needs exactly two new lines:
class SelfAttention(nn.Module):
def __init__(self, cfg):
super().__init__()
self.heads = cfg.heads
self.modern = cfg.modern
self.qkv = nn.Linear(cfg.dim, 3 * cfg.dim) # Q, K and V in one go
self.out = nn.Linear(cfg.dim, cfg.dim) # mixes the heads back together
self.drop = nn.Dropout(cfg.dropout)
mask = torch.tril(torch.ones(cfg.context, cfg.context, dtype=torch.bool))
self.register_buffer("mask", mask) # causal mask, not a parameter
def forward(self, x):
B, T, D = x.shape
d = D // self.heads
q, k, v = self.qkv(x).split(D, dim=-1) # each [B, T, D]
# split into heads: [B, T, D] -> [B, H, T, d]
q = q.view(B, T, self.heads, d).transpose(1, 2)
k = k.view(B, T, self.heads, d).transpose(1, 2)
v = v.view(B, T, self.heads, d).transpose(1, 2)
if self.modern:
q, k = rope(q), rope(k) # position by rotating instead of a table
scores = q @ k.transpose(-2, -1) / math.sqrt(d) # compare: [B, H, T, T]
scores = scores.masked_fill(~self.mask[:T, :T], float("-inf")) # block the future
weights = scores.softmax(dim=-1) # percentages per row
y = weights @ v # mix: [B, H, T, d]
y = y.transpose(1, 2).contiguous().view(B, T, D) # heads side by side again
return self.drop(self.out(y))And because the position now lives in the rotation, the modern GPT doesn't need a position table anymore.
SwiGLU: the MLP with a bouncer
class MLP(nn.Module):
def __init__(self, cfg):
super().__init__()
self.modern = cfg.modern
hidden = 4 * cfg.dim if not cfg.modern else 8 * cfg.dim // 3
self.up = nn.Linear(cfg.dim, hidden)
if cfg.modern:
self.gate = nn.Linear(cfg.dim, hidden) # SwiGLU: the bouncer
self.down = nn.Linear(hidden, cfg.dim)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x):
if self.modern:
return self.drop(self.down(F.silu(self.gate(x)) * self.up(x)))
return self.drop(self.down(F.gelu(self.up(x)))) # wider, kink, narrow againIn modern mode there is a second path: gate, the bouncer. F.silu(self.gate(x)) decides position by position how much of self.up(x) gets through. Because there are now three weight tables instead of two, we make the intermediate width smaller: 8D/3 instead of 4D. That keeps the parameter count comparable.
Block and GPT
class Block(nn.Module):
def __init__(self, cfg):
super().__init__()
Norm = RMSNorm if cfg.modern else nn.LayerNorm
self.norm1 = Norm(cfg.dim)
self.attn = SelfAttention(cfg)
self.norm2 = Norm(cfg.dim)
self.mlp = MLP(cfg)
def forward(self, x):
x = x + self.attn(self.norm1(x)) # positions exchange information
x = x + self.mlp(self.norm2(x)) # each position computes on its own
return xclass GPT(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.tokens = nn.Embedding(cfg.vocab_size, cfg.dim) # what
self.positions = None if cfg.modern else nn.Embedding(cfg.context, cfg.dim) # where (modern: RoPE)
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
self.norm = RMSNorm(cfg.dim) if cfg.modern else nn.LayerNorm(cfg.dim)
self.head = nn.Linear(cfg.dim, cfg.vocab_size, bias=False)
self.head.weight = self.tokens.weight # weight tying: one table, two jobs
self.apply(self.init_weights)
@staticmethod
def init_weights(module):
# small random starting values (chapter 20) – otherwise the first logits are huge
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, ids):
B, T = ids.shape
x = self.tokens(ids)
if self.positions is not None:
x = x + self.positions(torch.arange(T, device=ids.device))
for block in self.blocks:
x = block(x)
return self.head(self.norm(x)) # Logits: [B, T, V]In the block the switch only decides which normalization is used. In the GPT, modern=True drops the position table (positions = None), and forward only adds it if it exists.
The whole file
This is the final version of model.py. It can do both: classic and modern.
# A GPT from scratch: embedding, attention, MLP, blocks (chapters 15–19)
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
@dataclass
class Config:
vocab_size: int
context: int = 128 # T: how many tokens the model sees at once
dim: int = 128 # D: numbers per token
heads: int = 4 # H: attention heads
layers: int = 4 # L: transformer blocks
dropout: float = 0.2 # against memorizing (chapter 23)
modern: bool = False # RMSNorm + RoPE + SwiGLU instead of LayerNorm + position table + GELU (chapter 25)
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight
def rope(x):
"""Rotate pairs of numbers depending on position (x: [B, H, T, d])."""
T, d = x.shape[-2], x.shape[-1]
speed = 10000 ** (-torch.arange(0, d, 2, device=x.device) / d) # every pair turns at its own speed
angle = torch.arange(T, device=x.device)[:, None] * speed[None, :] # [T, d/2]
cos, sin = angle.cos(), angle.sin()
even, odd = x[..., 0::2], x[..., 1::2]
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)
class SelfAttention(nn.Module):
def __init__(self, cfg):
super().__init__()
self.heads = cfg.heads
self.modern = cfg.modern
self.qkv = nn.Linear(cfg.dim, 3 * cfg.dim) # Q, K and V in one go
self.out = nn.Linear(cfg.dim, cfg.dim) # mixes the heads back together
self.drop = nn.Dropout(cfg.dropout)
mask = torch.tril(torch.ones(cfg.context, cfg.context, dtype=torch.bool))
self.register_buffer("mask", mask) # causal mask, not a parameter
def forward(self, x):
B, T, D = x.shape
d = D // self.heads
q, k, v = self.qkv(x).split(D, dim=-1) # each [B, T, D]
# split into heads: [B, T, D] -> [B, H, T, d]
q = q.view(B, T, self.heads, d).transpose(1, 2)
k = k.view(B, T, self.heads, d).transpose(1, 2)
v = v.view(B, T, self.heads, d).transpose(1, 2)
if self.modern:
q, k = rope(q), rope(k) # position by rotating instead of a table
scores = q @ k.transpose(-2, -1) / math.sqrt(d) # compare: [B, H, T, T]
scores = scores.masked_fill(~self.mask[:T, :T], float("-inf")) # block the future
weights = scores.softmax(dim=-1) # percentages per row
y = weights @ v # mix: [B, H, T, d]
y = y.transpose(1, 2).contiguous().view(B, T, D) # heads side by side again
return self.drop(self.out(y))
class MLP(nn.Module):
def __init__(self, cfg):
super().__init__()
self.modern = cfg.modern
hidden = 4 * cfg.dim if not cfg.modern else 8 * cfg.dim // 3
self.up = nn.Linear(cfg.dim, hidden)
if cfg.modern:
self.gate = nn.Linear(cfg.dim, hidden) # SwiGLU: the bouncer
self.down = nn.Linear(hidden, cfg.dim)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x):
if self.modern:
return self.drop(self.down(F.silu(self.gate(x)) * self.up(x)))
return self.drop(self.down(F.gelu(self.up(x)))) # wider, kink, narrow again
class Block(nn.Module):
def __init__(self, cfg):
super().__init__()
Norm = RMSNorm if cfg.modern else nn.LayerNorm
self.norm1 = Norm(cfg.dim)
self.attn = SelfAttention(cfg)
self.norm2 = Norm(cfg.dim)
self.mlp = MLP(cfg)
def forward(self, x):
x = x + self.attn(self.norm1(x)) # positions exchange information
x = x + self.mlp(self.norm2(x)) # each position computes on its own
return x
class GPT(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.tokens = nn.Embedding(cfg.vocab_size, cfg.dim) # what
self.positions = None if cfg.modern else nn.Embedding(cfg.context, cfg.dim) # where (modern: RoPE)
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
self.norm = RMSNorm(cfg.dim) if cfg.modern else nn.LayerNorm(cfg.dim)
self.head = nn.Linear(cfg.dim, cfg.vocab_size, bias=False)
self.head.weight = self.tokens.weight # weight tying: one table, two jobs
self.apply(self.init_weights)
@staticmethod
def init_weights(module):
# small random starting values (chapter 20) – otherwise the first logits are huge
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, ids):
B, T = ids.shape
x = self.tokens(ids)
if self.positions is not None:
x = x + self.positions(torch.arange(T, device=ids.device))
for block in self.blocks:
x = block(x)
return self.head(self.norm(x)) # Logits: [B, T, V]Run and compare
First check that the classic model still works exactly the same:
python check.pyThe output must be the same as in step 4. Then switch on the modern variant in train.py. Change a single line:
cfg = Config(vocab_size=tok.vocab_size, modern=True)python train.py923,432 parameters on cpu
...
Best validation loss 2.340, saved to data/model.ptThe direct comparison on the same data, with the same settings:
| Classic | Modern | |
|---|---|---|
| Parameters | 940,800 | 923,432 |
| Best validation loss | 2.548 | 2.340 |
| Training time (on my machine) | 99 s | 146 s |
The modern model is better, and with fewer parameters. A loss difference of 0.2 sounds small but is clear: on average the model gives the correct next token about 23 percent more probability (e to the power of 0.21 ≈ 1.23). In exchange each step takes a bit longer, mostly because of the rotations.
This is exactly how research works: a hypothesis (“modern components help”), change one thing, measure, compare honestly. If you want to know precisely, switch on the three components one at a time; that would be a real ablation from chapter 25.
After training with modern=True, generate.py, sft.py and chat.py work without changes: they read the settings from the saved model.
Where to go from here
Your model now has the same components as many current open-source models. What it mostly lacks is size and data. Ideas for the next steps are in the following chapter, for example:
- GQA from chapter 26: several query heads share keys and values.
- A KV cache for faster generation.
- More data and a bigger model, ideally on a graphics card.
Key points
- A switch
modernswaps three components: RMSNorm instead of LayerNorm, RoPE instead of a position table, SwiGLU instead of the GELU MLP. - RoPE only rotates query and key. For that, attention needs two more lines, and the position table goes away.
- On our stories the modern model is measurably better, with fewer parameters.