Inside a Transformer Block

Building a tiny GPT from scratch

Bolt every piece from this section together into a working character-level language model that trains on a laptop CPU in under twenty seconds.

On this page 6
  1. What you are building
  2. What it will and will not do
  3. Why memorising is still worth watching
  4. What to do after it works
  5. Remember this
  6. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Everything in this section fits into about sixty lines of code. That code trains a real, working language model on an ordinary laptop.

Think about learning to fix a cycle. Somebody shows you the chain, then the brakes, then the gears, one at a time. You nod along. Then one day you bolt the whole thing together yourself and ride it down the lane.

That ride is different from all the explanations put together. It is small, it wobbles, and it is genuinely yours.

That is this lesson. A small model, on a tiny piece of text, but every part is the real thing.

What you are building

A model that reads characters and guesses the next one.

Not words — individual letters, spaces and full stops. That keeps the vocabulary tiny, around twenty different characters. The whole thing then trains in seconds instead of days.

Everything else is genuine. Attention with a mask. Blocks stacked with additions. Normalisation before each half. A shared table for input and output. Next-character prediction as the training task. Make it bigger and give it more text and it becomes a real language model.

What it will and will not do

It will learn the text you give it. Show it five sentences repeated, and it learns those five sentences, including the order they come in.

It will not say anything new. There is nothing else in there. It has a hundred and fifty thousand parameters and a few hundred characters of text. Memorising is the only thing available to it.

Say that plainly to yourself before you run it. A model this size memorises. It is not a small ChatGPT. It is the same machinery at a scale where memorising is what the machinery does.

Why memorising is still worth watching

Because you can see the whole arc, live, in fifteen seconds.

At the start the model is guessing randomly. Its error is exactly what random guessing over twenty-two characters should give. Within a hundred steps it has dropped near zero. Then you type in a few characters and it continues the sentence correctly.

That whole loop — random, then learning, then working — is what training is. Seeing it happen on your own machine, with numbers you can check, beats reading about it.

What to do after it works

Change one thing at a time and watch what happens.

  • Remove the mask and see the error drop implausibly fast, because the model is now reading its own answer.
  • Remove the additions between blocks and watch training get much worse.
  • Give it more text than it can memorise and watch the error stop going to zero.

Each of those is a lesson from this section, felt rather than read.

Remember this

  • A working language model is around sixty lines and trains on a laptop in seconds.
  • At this size it memorises, which is honest and still worth seeing.
  • Breaking it deliberately, one piece at a time, is the fastest way to understand it.

What to learn next

  • Digit recognition — another end-to-end build, on images instead of text.
  • How LLMs work — what changes when this is scaled up a hundred thousand times.
  • Fine-tuning — taking a trained model and adapting it, rather than starting from nothing.

Developer — Code and libraries.

Setup

bash
pip install torch

No dataset download. The training text is inline, and the whole run finishes in under twenty seconds on a laptop CPU.

The complete model

tiny_gpt.py
import time
import torch
import torch.nn as nn
import torch.nn.functional as F

TEXT = ("pranay makes chai every morning. "
        "pranay makes chai for his friends. "
        "the chai is hot and sweet. "
        "the chai is ready every morning. "
        "his friends drink the chai and smile. ") * 12

chars = sorted(set(TEXT))
stoi = {c: i for i, c in enumerate(chars)}
itos = {i: c for c, i in stoi.items()}
V = len(chars)
data = torch.tensor([stoi[c] for c in TEXT])

D, L, H, BLOCK = 64, 3, 4, 32

class Block(nn.Module):
    def __init__(self):
        super().__init__()
        self.n1, self.n2 = nn.LayerNorm(D), nn.LayerNorm(D)
        self.qkv, self.proj = nn.Linear(D, 3 * D), nn.Linear(D, D)
        self.f1, self.f2 = nn.Linear(D, 4 * D), nn.Linear(4 * D, D)
    def forward(self, x):
        B, T, _ = x.shape
        q, k, v = self.qkv(self.n1(x)).chunk(3, -1)
        q, k, v = (z.view(B, T, H, D // H).transpose(1, 2) for z in (q, k, v))
        a = F.scaled_dot_product_attention(q, k, v, is_causal=True)   # the mask lives here
        x = x + self.proj(a.transpose(1, 2).reshape(B, T, D))
        return x + self.f2(F.gelu(self.f1(self.n2(x))))

class GPT(nn.Module):
    def __init__(self):
        super().__init__()
        self.tok, self.pos = nn.Embedding(V, D), nn.Embedding(BLOCK, D)
        self.blocks = nn.Sequential(*[Block() for _ in range(L)])
        self.norm = nn.LayerNorm(D)
        self.head = nn.Linear(D, V, bias=False)
        self.head.weight = self.tok.weight                           # tied embeddings
        self.apply(self._init)
    @staticmethod
    def _init(m):
        if isinstance(m, (nn.Linear, nn.Embedding)):
            nn.init.normal_(m.weight, std=0.02)          # GPT-2's init; the default is far too big
            if getattr(m, "bias", None) is not None:
                nn.init.zeros_(m.bias)
    def forward(self, idx):
        h = self.tok(idx) + self.pos(torch.arange(idx.size(1)))
        return self.head(self.norm(self.blocks(h)))

torch.manual_seed(1337)
model = GPT()
print(f"vocabulary {V} characters, {sum(p.numel() for p in {id(p): p for p in model.parameters()}.values()):,} parameters")

opt = torch.optim.AdamW(model.parameters(), lr=3e-3)
g = torch.Generator().manual_seed(0)
t0 = time.time()
for step in range(601):
    i = torch.randint(len(data) - BLOCK - 1, (32,), generator=g)
    x = torch.stack([data[j:j + BLOCK] for j in i])
    y = torch.stack([data[j + 1:j + BLOCK + 1] for j in i])
    loss = F.cross_entropy(model(x).reshape(-1, V), y.reshape(-1))
    opt.zero_grad(); loss.backward(); opt.step()
    if step % 150 == 0:
        print(f"  step {step:>3}  loss {loss.item():.4f}")
print(f"trained in {time.time() - t0:.0f} seconds on a laptop CPU")

@torch.no_grad()
def generate(prompt, n=90):
    idx = torch.tensor([[stoi[c] for c in prompt]])
    for _ in range(n):
        logits = model(idx[:, -BLOCK:])[:, -1]
        nxt = logits.argmax(-1, keepdim=True)          # greedy: always the top guess
        idx = torch.cat([idx, nxt], dim=1)
    return "".join(itos[int(i)] for i in idx[0])

print("\ngreedy continuation of 'pranay '")
print(" ", generate("pranay "))
print("\ngreedy continuation of 'the chai '")
print(" ", generate("the chai "))
Output
vocabulary 22 characters, 153,536 parameters
  step   0  loss 3.1352
  step 150  loss 0.1059
  step 300  loss 0.0918
  step 450  loss 0.0888
  step 600  loss 0.0884
trained in 12 seconds on a laptop CPU

greedy continuation of 'pranay '
  pranay makes chai every morning. pranay makes chai for his friends. the chai is hot and sweet. th

greedy continuation of 'the chai '
  the chai is ready every morning. his friends drink the chai and smile. pranay makes chai every morn

Run this on PyTorch 2.5.1 with these seeds and the loss values reproduce exactly. The elapsed seconds depend on your CPU. A different PyTorch version or platform can shift the last decimal places of the loss. At these tiny margins, that can occasionally change a generated character.

The first loss value is a free correctness check

3.1352. The natural logarithm of 22 is 3.0910.

A model that has learned nothing assigns equal probability to all 22 characters. The cross-entropy is then exactly ln(22). Starting within one percent of that means the initialisation is sane and nothing is broken before training begins.

This check costs nothing and catches a great deal. If your first loss is 46 instead of 3.1, your weights are initialised too large. That is exactly what happens without the _init method. nn.Embedding's default fills the tied matrix from a standard normal distribution. If your first loss is far below ln(V), something is leaking the answer.

Where every piece of this section shows up

LineLesson
self.tok, self.postokens and positions become vectors
self.n1(x) before the attentionpre-norm placement
self.qkv(...).chunk(3, -1)queries, keys and values
.view(B, T, H, D // H).transpose(1, 2)multi-head attention
is_causal=Truecausal masking
x = x + ... twicethe residual stream
self.f2(F.gelu(self.f1(...)))the feedforward layer
self.norm before self.headthe final norm a pre-norm model needs
self.head.weight = self.tok.weighttied embeddings

Nothing has been simplified away. Scale D, L and the text up. Replace the learned position table with rotary positions. Swap LayerNorm for RMSNorm and the feedforward layer for SwiGLU. That is a current-generation architecture.

Details in the code that are easy to get wrong

idx[:, -BLOCK:] in the generation loop. The position table has only BLOCK rows. Feed a longer sequence and you get an index error. Real models solve this with a much larger table or with rotary positions, which have no table at all.

The deduplication in the parameter count. {id(p): p for p in model.parameters()} is needed here. The tied matrix would otherwise be counted twice. PyTorch's own parameters() already deduplicates, so the plain sum gives the same answer here. The explicit version is written out so the reason is visible.

@torch.no_grad() on generate. Without it, every generated token extends the autograd graph and memory grows until the process dies. This is one of the most common bugs in hand-written generation loops.

torch.arange(idx.size(1)) recomputed each call. Correct but wasteful. A real implementation registers it as a buffer. See buffers vs parameters.

Four experiments, in order of what they teach

Remove the mask. Change is_causal=True to False. The loss collapses toward zero far faster, and generation becomes gibberish. The model has learned to copy the answer sitting to its right. That answer is present during training and absent during generation.

Remove the residual additions. Change both x = x + ... lines to x = .... Training gets slower and worse. At three layers the damage is modest; the point is the direction, and it grows with depth.

Remove the activation. Delete F.gelu. Two stacked linear layers collapse into one, and the feedforward half of every block becomes a single matrix. Loss stalls at a visibly higher value.

Give it more text than it can hold. Paste in a few pages of your own writing. The loss stops going near zero and settles somewhere higher. That is the first honest language model you will have trained. It can no longer memorise its way out.

Sampling instead of greedy decoding

argmax always takes the top guess, which is why the output above is a clean recital. For varied output, sample from the distribution instead:

python
probs = (logits / temperature).softmax(-1)
nxt = torch.multinomial(probs, num_samples=1)

No output block for this one, deliberately. It is random by design. A printed sample would teach you to expect something that will not happen. Higher temperature spreads the probability out and produces more variety and more mistakes. See temperature and sampling.

Common mistakes

Forgetting opt.zero_grad(). Gradients accumulate across steps by default. The update grows without bound, and the loss goes to nan within a few steps.

Shifting the target incorrectly. y must be x moved one position left. Off by one in either direction gives a model that trains and generates nonsense.

Leaving the final self.norm out. A pre-norm stack has nothing rescaling its output. Skipping the last norm produces oddly saturated logits.

Training on a CPU and expecting it to scale. This runs in seconds because the model is tiny. Multiply D by ten and you will want a GPU. See installing PyTorch with CUDA.

Try it yourself

Run the parameter formula from counting a model's parameters by hand on this configuration. Use V=22, d=64, L=3, d_ff=256, max_pos=32, tied. Check it against the printed 153,536. If it matches, you can size any transformer from its config file. If it does not, the difference tells you which component you got wrong.

What to learn next

  • Digit recognition — another end-to-end build, on images instead of text.
  • How LLMs work — what changes when this is scaled up a hundred thousand times.
  • Fine-tuning — taking a trained model and adapting it, rather than starting from nothing.

Researcher — Mathematics and papers.

What this implementation is, precisely

A GPT-2-shaped decoder. Pre-norm LayerNorm, learned absolute position embeddings, a GELU feedforward with a four-times multiplier. Tied input and output embeddings, causal self-attention, next-token cross-entropy. The differences from GPT-2 small are scale only: $d = 64$ against 768, $L = 3$ against 12, $|V| = 22$ against 50,257. There is also a character-level vocabulary in place of byte-pair encoding.

Parameter count from the standard formula, with $\tau = 1$ for tying:

$$ N = |V| d + T_{\max} d + L(12 d^2 + 9d) + 2d $$

$= 22 \cdot 64 + 32 \cdot 64 + 3(12 \cdot 4096 + 576) + 128 = 153{,}536$, matching the printed count exactly. The $9d$ term collects the biases on the four linear layers. It also collects the gains and biases of the two norms.

Initialisation, and why the first loss is diagnostic

At initialisation, a well-conditioned model outputs near-uniform logits, giving

$$ \mathcal{L}_0 \approx \ln |V| $$

Here $\ln 22 = 3.0910$ against a measured 3.1352. This is the cheapest available sanity check on a language model and it is under-used.

The std=0.02 initialisation is GPT-2's. It matters more than usual here because of weight tying. nn.Embedding defaults to $\mathcal{N}(0,1)$. A tied head with unit-variance rows produces logits of standard deviation $\sqrt{d} \approx 8$, giving an initial loss around 46. Correcting that costs several hundred optimisation steps that do nothing but shrink the output scale.

GPT-2 additionally scales residual-projection weights by $1/\sqrt{2L}$ to keep the residual stream's variance stable with depth. At $L=3$ the effect is negligible; at $L=48$ it is not.

What a model at this scale can and cannot represent

Training corpus: 1,992 characters, of which only 166 are distinct before the twelvefold repetition. Model: 153,536 parameters. The ratio makes memorisation the optimal solution by a wide margin. The loss plateau near 0.088 rather than 0.0 reflects genuine ambiguity in the corpus. After "the chai is " the text continues with either "hot" or "ready". No context inside a 32-character window disambiguates them.

That plateau is worth noting because it is not a training failure. It is the conditional entropy of the data given the model's context. A perfect model would reach it too. Distinguishing an irreducible floor from a fixable one is a core skill in reading a loss curve.

Scaling this to something real

The nanoGPT reference implementation (Karpathy) is essentially this file, plus the additions that matter at scale. It is the right next step:

  • Byte-pair tokenization rather than characters. The vocabulary goes from tens to tens of thousands, making the embedding table a significant share of parameters.
  • Mixed precision with bfloat16 autocast, plus gradient clipping.
  • Cosine learning-rate decay with warm-up. Weight decay applied to matrix parameters, but not to norms or biases.
  • Gradient accumulation to reach a large effective batch, and torch.compile.
  • A held-out validation split, without which loss reaching zero tells you nothing.

The last point is the important one pedagogically. This script has no validation split. That is acceptable for a demonstration whose purpose is to memorise, and dishonest in anything else.

Modernising the architecture

Four substitutions convert this into a current-generation decoder:

ReplaceWithReference
learned position embeddingsrotary position embeddingsSu et al., 2021, arXiv:2104.09864
LayerNormRMSNormZhang and Sennrich, 2019, arXiv:1910.07467
GELU feedforwardSwiGLU at $\tfrac{8}{3}d$Shazeer, 2020, arXiv:2002.05202
full multi-head attentiongrouped-query attentionAinslie et al., 2023, arXiv:2305.13245

Each is a local change of a few lines. None alters the training loop, the loss, or the residual structure. The architecture has been this stable for eight years. It absorbed four substantial component swaps without a change to its skeleton. That is the most interesting fact about it.

Papers

What to learn next

  • Digit recognition — another end-to-end build, on images instead of text.
  • How LLMs work — what changes when this is scaled up a hundred thousand times.
  • Fine-tuning — taking a trained model and adapting it, rather than starting from nothing.