Tokenwerk · The LLM textbook

Chapter 36 · VII · Build your LLM · 12 minutes

Step 4: The GPT model

Attention with mask and heads, MLP, blocks, embeddings: the complete GPT in about 90 lines, and three checks before a single second of training.

What we build

Now for the heart of it: a real GPT. It is exactly the decoder from chapter 19 (embeddings, attention, MLP, residual connections, normalization), only this time as code. Everything goes into the file model.py, a little over 90 lines.

The blueprint from the inside out:

  1. Config: all sizes in one place.
  2. SelfAttention: positions exchange information (chapters 17 and 18).
  3. MLP: each position keeps computing on its own.
  4. Block: attention and MLP plugged together with normalization and residuals.
  5. GPT: embeddings, a stack of blocks, output head.

At the end we check with check.py that everything is wired correctly, before training for a single second.

The settings

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

A dataclass is a class that only stores values. That way every size appears exactly once, and all building blocks read it from cfg. You know the letters from chapter 11: T, D, H and L. With these values our model gets just under a million parameters: tiny compared with ChatGPT, but enough to learn simple stories.

dropout is new: during training, a fifth of the values are randomly set to zero. That forces the model not to rely on single connections and slows down memorizing. Step 5 shows when this matters.

Self-attention

python
class SelfAttention(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.heads = cfg.heads
        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)
        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))

This is chapters 17 and 18 in 25 lines. Let's go through forward line by line:

  • Compute Q, K, V: self.qkv is a single linear layer that turns every list of width D directly into three lists: query, key and value. That is the same as three separate layers, just faster. split(D, dim=-1) cuts the result back into three parts.
  • Split into heads: view(B, T, H, d) cuts every list into H pieces of width d = D/H. transpose(1, 2) moves the head axis forward: [B, T, H, d] becomes [B, H, T, d]. Now each head computes on its own (chapter 18).
  • Compare: q @ k.transpose(-2, -1) computes all dot products between queries and keys at once: the score matrix of shape [B, H, T, T]. Divided by √d, as in chapter 17.
  • Block the future: self.mask is a triangle of True/False. Wherever it says False (above and to the right of the diagonal), masked_fill sets the score to −∞. register_buffer means: the mask belongs to the model but isn't trained.
  • Weight and mix: softmax turns every row into percentages, weights @ v mixes the values.
  • Put the heads back together: transpose back, all heads side by side again, then self.out mixes the results.

The MLP

python
class MLP(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.up = nn.Linear(cfg.dim, 4 * cfg.dim)
        self.down = nn.Linear(4 * cfg.dim, cfg.dim)
        self.drop = nn.Dropout(cfg.dropout)

    def forward(self, x):
        return self.drop(self.down(F.gelu(self.up(x))))          # wider, kink, narrow again

Make it wider (D → 4D), kink (GELU), narrow again (4D → D). Exactly as in chapter 19.

The block

python
class Block(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.norm1 = nn.LayerNorm(cfg.dim)
        self.attn = SelfAttention(cfg)
        self.norm2 = nn.LayerNorm(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

Two lines in forward, and yet the whole transformer block: normalize, attention, add the result. Then normalize, MLP, add again. The x = x + ... is the residual connection: add instead of replace.

The whole GPT

python
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 = nn.Embedding(cfg.context, cfg.dim)      # where
        self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
        self.norm = 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) + 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]
  • tokens and positions are two embedding tables: one for what (which token), one for where (which position). Both lists are added.
  • nn.ModuleList holds the stack of blocks. A normal Python list wouldn't work, because then PyTorch wouldn't find the weights inside it (chapter 20).
  • self.head.weight = self.tokens.weight is weight tying from chapter 15: the input and output tables share the same numbers.
  • init_weights gives all weights small starting values. That is not a detail! Without this line, training starts at a loss of 87 instead of 6.9. The reason: PyTorch fills the embedding table with large random numbers, and thanks to weight tying they end up in the output head too. Huge logits mean a model that is very sure, and very wrong.
  • forward tells the whole story from chapter 19 in five lines: look up, add the position, through all blocks, normalize, output head.

The whole file

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


class SelfAttention(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.heads = cfg.heads
        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)
        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.up = nn.Linear(cfg.dim, 4 * cfg.dim)
        self.down = nn.Linear(4 * cfg.dim, cfg.dim)
        self.drop = nn.Dropout(cfg.dropout)

    def forward(self, 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__()
        self.norm1 = nn.LayerNorm(cfg.dim)
        self.attn = SelfAttention(cfg)
        self.norm2 = nn.LayerNorm(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 = nn.Embedding(cfg.context, cfg.dim)      # where
        self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])
        self.norm = 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) + 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]

Check before we train

A bug in the model often only shows up after hours of training, or never, because the loss goes down anyway. So we check three things first (chapters 18 and 20). Create the file check.py:

python
# Checks before we train (chapters 18 and 20)
import torch
import torch.nn.functional as F

from model import GPT, Config

torch.manual_seed(0)
cfg = Config(vocab_size=1024, dropout=0.0)
model = GPT(cfg).eval()

ids = torch.randint(0, cfg.vocab_size, (2, 16))
logits = model(ids)
print("Output shape:", tuple(logits.shape), "(expected: (2, 16, 1024))")
print("Parameters:", f"{sum(p.numel() for p in model.parameters()):,}")

# 1. Our attention must give the same result as PyTorch's built-in version
attn = model.blocks[0].attn
x = torch.randn(2, 16, cfg.dim)
q, k, v = attn.qkv(x).split(cfg.dim, dim=-1)
split = lambda t: t.view(2, 16, cfg.heads, -1).transpose(1, 2)
reference = F.scaled_dot_product_attention(split(q), split(k), split(v), is_causal=True)
reference = attn.out(reference.transpose(1, 2).reshape(2, 16, cfg.dim))
print("Same as PyTorch:", torch.allclose(attn(x), reference, atol=1e-5))

# 2. Peeking test: a different ending must not change earlier predictions
a = torch.tensor([[1, 5, 6, 7, 8, 9]])
b = torch.tensor([[1, 5, 6, 7, 8, 3]])
same = torch.allclose(model(a)[:, :5], model(b)[:, :5], atol=1e-5)
print("Doesn't look into the future:", same)

# 3. The starting loss should be about ln(1024) = 6.93
targets = torch.randint(0, cfg.vocab_size, (2, 16))
loss = F.cross_entropy(logits.view(-1, cfg.vocab_size), targets.view(-1))
print(f"Starting loss: {loss.item():.2f}")
  1. Is our attention right? PyTorch has a built-in, heavily optimized attention function. If our own calculation gives the same result, we didn't make a calculation error.
  2. Does the model look into the future? The peeking test from chapter 18: two sequences that only differ at the very end must have the same predictions before that.
  3. Does the model start neutral? An untrained model shouldn't have an opinion. Then the error is about ln(1024) ≈ 6.93, as if it were guessing evenly among all 1024 tokens.

.eval() and dropout=0.0 switch off the randomness of dropout, so the comparisons are exact.

bash
python check.py
text
Output shape: (2, 16, 1024) (expected: (2, 16, 1024))
Parameters: 940,800
Same as PyTorch: True
Doesn't look into the future: True
Starting loss: 6.93

Three green lights. Your GPT is built correctly.

If something goes wrong

  • Same as PyTorch: False: usually a mistake in the mask or in dividing by √d. Compare the lines with scores carefully with the book.
  • Doesn't look into the future: False: the mask is missing or the wrong way round. torch.tril (lower triangle) is right; torch.triu would be wrong.
  • Starting loss clearly above 7: self.apply(self.init_weights) is missing.
  • RuntimeError: view size is not compatible: you forgot .contiguous() before the last view.

Key points

  • model.py rebuilds chapters 15–19: embeddings, self-attention with mask and heads, MLP, block with residual, output head with weight tying.
  • Small starting weights matter: otherwise training begins with a huge error.
  • check.py checks before training: attention computes correctly, no look into the future, a neutral start at 6.93.