Tokenwerk · The LLM textbook

Chapter 40 · VII · Build your LLM · 12 minutes

Step 8: Modernizing

Build RMSNorm, RoPE and SwiGLU into your model, and measure whether they really help. Spoiler: yes, with fewer parameters.

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

python
@dataclass

RMSNorm

python
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

Like 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

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

  • speed sets how fast each pair of numbers turns: the first pair like a second hand, the last like an hour hand.
  • angle is position times speed: at position 0 nothing turns, at position 10 it turns ten times as far as at position 1.
  • even and odd are the two numbers of each pair. The line with cos and sin is the rotation formula from the math corner of chapter 25.
  • stack and flatten put 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:

python
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

python
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

In 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

python
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
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 = 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.

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

bash
python check.py

The output must be the same as in step 4. Then switch on the modern variant in train.py. Change a single line:

python
cfg = Config(vocab_size=tok.vocab_size, modern=True)
bash
python train.py
text
923,432 parameters on cpu
...
Best validation loss 2.340, saved to data/model.pt

The 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 modern swaps 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.