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.
- 14 min read
- 3 reading levels
- Updated
On this page 6
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
pip install torchNo dataset download. The training text is inline, and the whole run finishes in under twenty seconds on a laptop CPU.
The complete model
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 "))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
| Line | Lesson |
|---|---|
self.tok, self.pos | tokens and positions become vectors |
self.n1(x) before the attention | pre-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=True | causal masking |
x = x + ... twice | the residual stream |
self.f2(F.gelu(self.f1(...))) | the feedforward layer |
self.norm before self.head | the final norm a pre-norm model needs |
self.head.weight = self.tok.weight | tied 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:
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:
| Replace | With | Reference |
|---|---|---|
| learned position embeddings | rotary position embeddings | Su et al., 2021, arXiv:2104.09864 |
| LayerNorm | RMSNorm | Zhang and Sennrich, 2019, arXiv:1910.07467 |
| GELU feedforward | SwiGLU at $\tfrac{8}{3}d$ | Shazeer, 2020, arXiv:2002.05202 |
| full multi-head attention | grouped-query attention | Ainslie 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
- Radford et al., Language Models are Unsupervised Multitask Learners (GPT-2), 2019
- Vaswani et al., Attention Is All You Need, 2017 — arxiv.org/abs/1706.03762
- Kaplan et al., Scaling Laws for Neural Language Models, 2020 — arxiv.org/abs/2001.08361
- Hoffmann et al., Training Compute-Optimal Large Language Models, 2022 — arxiv.org/abs/2203.15556
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.