Tokenwerk · The LLM textbook

Chapter 20 · III · Building the transformer · 4 minutes

Implementing your own decoder

From formulas to a model class. A complete, runnable reference implementation backs every decision.

The project is your runnable blueprint

In the download you'll find model.py. The file contains small, separate building blocks (classes and a RoPE function) for configuration, normalization, position rotation, attention, feed-forward, block and the full model. It does not load any model weights from the internet. The two variants differ in clearly named building blocks; the training objective stays the same.

Read the file from the bottom up first: the top-level class shows which modules are put together. Then go into the block, and only after that into the attention details. Reading in this order keeps the main data path in view.

A minimal configuration

python
from model import Config, LM
cfg = Config(
    vocab_size=320,
    context=128,
    dim=64,
    layers=2,
    heads=4,
    kv_heads=4,
    modern=False,
)
model = LM(cfg)
print(sum(p.numel() for p in model.parameters()))

This is a very small debugging configuration. dim must be divisible by heads. For the modern RoPE variant, the head dimension must be even. With GQA, the number of query heads must be divisible by the number of key/value heads. The configuration check catches these cases before the first training run.

After data preparation, the vocabulary size comes from the tokenizer that was actually learned. If it found fewer merges than requested, you use its actual size. You must not just write a desired size into the output and thereby create invisible classes.

ModuleList instead of a plain list

python
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.layers)])

With ModuleList, PyTorch recognizes the contained modules, includes them in state_dict() and moves them with .to(device). A plain Python list does not register its contained modules in the same way. The forward pass might run, but model.parameters() would miss important weights.

"Using one block several times" and "creating several blocks" are also different things. [block] * L shares the same block object. Our code creates each block object separately. Sharing parameters between layers would be a possible different architecture, but not an accidental implementation detail.

Initialization

Our reference model initializes embeddings and linear weights from a small normal distribution, and biases to zero. Norm scales start at one. This choice is a clear starting point, not a rule that transfers optimally to every model. A larger model may need additional depth-dependent scaling.

Weight tying happens after the modules are initialized, so the tied matrix isn't accidentally initialized twice under conflicting assumptions. The optimizer receives the unique parameters from model.parameters(); the shared matrix is not trained as two independent weights.

Checking the forward pass

python
import torch
ids = torch.randint(0, cfg.vocab_size, (2, 16))
logits = model(ids)
assert logits.shape == (2, 16, cfg.vocab_size)
assert torch.isfinite(logits).all()

If you exceed the context length, the model should raise an understandable error message. A silent position overflow is not useful behavior. With too little input, generating a next token from an empty array is also problematic; BOS solves this case cleanly.

Checking the backward pass

python
import torch.nn.functional as F
x, y = ids[:, :-1], ids[:, 1:]
loss = F.cross_entropy(model(x).reshape(-1, cfg.vocab_size), y.reshape(-1))
loss.backward()
assert model.embed.weight.grad is not None
assert torch.isfinite(model.embed.weight.grad).all()

Don't check only the final output head. If a mistaken .detach() operation cuts the graph, early layers can remain untrained. The loss can still go down thanks to the later layers.

Train mode and eval mode

model.train() and model.eval() control mode-sensitive layers such as dropout. They do not switch automatic differentiation on or off. For that, you use torch.no_grad() or torch.inference_mode(). For validation, we set eval mode and turn off gradient recording.

Our small reference model has no dropout. The mode switches stay in the code anyway, because they are a correct basic pattern for later extensions.

Your first meaningful model test

With a fixed seed, compare the logits of two sequences with the same prefix and a different suffix. Check that the early logits stay the same. Repeat the test for the classic and the modern decoder. Then check that an optimizer step changes at least some weights.

A working forward pass and a falling test loss do not prove that your decoder is stable on large amounts of data. But they narrow down where bugs can hide considerably. Keep these small contract tests separate from long-term capability measurements.

Changing the architecture on purpose

At first, change only the number of blocks. Leave tokenizer, width, number of heads, context, data and token budget unchanged. Note the parameter count and the time per step. Observe that a deeper model has not only more capacity but also more compute time. A comparison after the same number of steps means something different from a comparison at an identical compute budget.