Skip to content
Go back

Train a Tiny GPT, Part 2: Attention

By KingPin 17 min read
Train a Tiny GPT, Part 2: Attention
Contents

Part 1 left us with one vector per token. Each of those vectors knows nothing about its neighbors. “Holmes” and ” telegram” sit in the same sentence and never talk to each other. A model that cannot connect them cannot do anything useful with language.

Attention is the mechanism that connects them, and the thesis fits in four facts. Attention is a weighted average where the weights are computed from the data. It is about 20 lines of code. It has no idea about word order until you add position information separately. And its cost grows with the square of the sequence length, which is why long context is expensive. This post builds it, breaks it, fixes it, and measures the bill on a CPU and an 8GB GPU.

Train a Tiny GPT at Home series: Part 1: Tokenizers · Part 2: Attention (you are here) · Part 3: Training · Part 4: Inference

Full example: Clone the working files at github.com/KingPin/sumguy-examples/…/part-2-attention

Run everything from llm/tiny-gpt/ after the setup in the series README. attention.py and position.py load the tokenizer.json that Part 1 produced, so train that tokenizer first.

Terminal window
cd part-2-attention
python attention.py
python position.py
python scaling.py # CPU, fp32
python scaling.py cuda # GPU, fp16

One Head, Spelled Out

Start with one attention head. Every token vector x gets multiplied by three matrices to produce three new vectors:

The score between token i and token j is the dot product of i’s query with j’s key. A big score means “j is relevant to i”. Softmax turns each row of scores into weights that sum to 1. The output for token i is the weighted sum of all the value vectors.

Here is the whole thing from attention.py:

attention.py
def one_head(x, wq, wk, wv):
"""x: (T, d). Returns (T, d_head) plus the attention weights (T, T)."""
q, k, v = x @ wq, x @ wk, x @ wv # each (T, d_head)
scores = q @ k.T / math.sqrt(q.shape[-1]) # (T, T): how much token i cares about token j
T = x.shape[0]
future = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)
scores = scores.masked_fill(future, float("-inf")) # no peeking ahead
weights = scores.softmax(dim=-1) # each row sums to 1
return weights @ v, weights

Two details need explaining.

The division by the square root of d_head. A dot product of two 64-dimensional vectors is a sum of 64 products. Its spread grows with the dimension. Large scores push softmax toward a one-hot spike, where one token gets a weight near 1 and the rest get nothing. Dividing by the square root of the head size keeps the scores in a range where softmax still has a useful gradient.

The causal mask. A GPT predicts the next token, so token i must never see token j when j is greater than i. Otherwise training is trivial: the model reads the answer. The triu call builds a boolean matrix that marks the upper triangle (the future). Filling those scores with minus infinity makes softmax give them a weight of exactly zero.

Run it on “Holmes examined the telegram.” and you get this:

5 tokens: ['Holmes', ' examined', ' the', ' telegram', '.']
single head matches F.scaled_dot_product_attention
attention weights (row = token doing the looking, untrained):
'Holmes' ' examined' ' the' ' telegram' '.'
'Holmes' 1.000 0.000 0.000 0.000 0.000
' examined' 0.626 0.374 0.000 0.000 0.000
' the' 0.302 0.272 0.426 0.000 0.000
' telegram' 0.352 0.399 0.063 0.187 0.000
'.' 0.587 0.044 0.292 0.033 0.044

Read it row by row. Each row is one token doing the looking. Row 1 is always 1.000, because the first token can only see itself. Every row sums to 1. The upper triangle is zero, which is the mask at work. The second line of output confirms the hand-written version matches PyTorch’s F.scaled_dot_product_attention to within 1e-5.

One warning before you read meaning into the numbers. The matrices here are random and untrained, so ”.”, looking at “Holmes” with weight 0.587, tells you nothing yet. Part 3 trains these matrices, and the weights start to mean something.

Many Heads, One Matmul

One head gives each token one way to look at the others. Real models run several heads in parallel, so one head can track, say, the nearest verb while another tracks the subject. The model splits its 256 channels into 4 heads of 64 dimensions each. Each head attends separately, and the results get glued back together.

attention.py
class MultiHeadAttention(nn.Module):
def __init__(self, d_model: int, n_heads: int):
super().__init__()
assert d_model % n_heads == 0
self.n_heads = n_heads
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False) # all heads, one matmul
self.out = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, fused: bool = True):
B, T, C = x.shape
q, k, v = self.qkv(x).split(C, dim=-1)
# (B, T, C) -> (B, heads, T, head_dim): each head gets its own slice of C
q, k, v = (t.view(B, T, self.n_heads, -1).transpose(1, 2) for t in (q, k, v))
if fused:
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
else:
scores = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1])
mask = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)
y = scores.masked_fill(mask, float("-inf")).softmax(-1) @ v
y = y.transpose(1, 2).reshape(B, T, C) # glue the heads back together
return self.out(y)

Two tricks keep this short.

The first is the single qkv layer. Instead of twelve small matrices (q, k and v for each of 4 heads), one linear layer produces all three tensors for all heads at once, and split cuts the result into q, k and v. One big matmul keeps the hardware busier than twelve small ones.

The second is the view and transpose pair. The view reshapes the 256 channels into 4 groups of 64. The transpose moves the head axis next to the batch axis, so every head becomes an independent attention problem. After attention, the reverse transpose and a reshape put the groups back in order.

The parameter count is easy to check by hand. The qkv layer is 3 x 256 x 256 = 196,608 weights. The out layer is 256 x 256 = 65,536. The sum is 4 x 256 x 256 = 262,144. Splitting into heads adds nothing. The head count only decides how the 256 channels get divided up.

multi-head (4 x 64): hand-rolled == fused, output (1, 5, 256), 262,144 params

The hand-rolled branch and the fused branch agree, and the script asserts it.

Attention Cannot Read

Look at the attention math again. Nothing in it mentions where a token sits in the sentence. Scores come from dot products of vectors, and the output is a weighted sum. If you shuffle the inputs, the same vectors get the same scores, and the outputs come out shuffled the same way. Order never enters the computation.

position.py proves it. With no mask and no position information, it permutes the five tokens and checks the output:

position.py
x = emb(torch.tensor(tok.encode("Holmes examined the telegram.")))
perm = torch.tensor([3, 0, 4, 1, 2])
a = attend(x, wq, wk, wv, causal=False, use_rope=False)
b = attend(x[perm], wq, wk, wv, causal=False, use_rope=False)
assert torch.allclose(a[perm], b, atol=1e-6)

That passes. A shuffle in means the same shuffle out.

A GPT uses the causal mask, though, and the mask changes the picture slightly. With the mask, each position sees a different number of tokens: position 1 sees one token, position 5 sees five. That count leaks a little order information into the model. So the honest claim is narrower than “masked attention is blind to order”. The test that matters is the last token. It sees every earlier token, so in a single layer the sentence is an unordered set from where it sits. Stack more layers and the earlier positions, which saw different counts, can pass order clues along. That is why real models still add position explicitly.

Take two sentences with opposite meanings: “Holmes shot Moriarty.” and “Moriarty shot Holmes.” Compare the final ”.” vector from each.

no position, no mask: shuffle [3, 0, 4, 1, 2] in -> same shuffle out
no position: max difference in final '.' vector = 0.000000

The difference is 0.000000. The final ”.” gets an identical vector for both sentences. The same four token vectors go into a softmax-weighted sum, and addition does not care about order. Whoever shot whom, the ”.” cannot tell in this one-layer test.

The fix is to inject position somewhere. There are two common ways. Learned absolute embeddings add a vector from a lookup table to each token, indexed by position, like Part 1’s token table but with positions as the row index. That table has a fixed number of rows, so the model cannot handle a position past its trained length. Rotary position embeddings, or RoPE, take a different route, and most current open models use RoPE or a variant of it.

RoPE: Position as a Rotation

RoPE does not add anything to the token vector. It rotates the query and key vectors by an angle that grows with position. The method comes from the RoFormer paper. Here is the implementation from position.py:

position.py
def rope(x, base: float = 10000.0):
"""Rotate pairs of channels in x (..., T, d) by position-dependent angles."""
T, d = x.shape[-2], x.shape[-1]
freqs = base ** (-torch.arange(0, d, 2, dtype=torch.float32) / d) # (d/2,) fast to slow
angles = torch.arange(T, dtype=torch.float32)[:, None] * freqs # (T, d/2)
cos, sin = angles.cos(), angles.sin()
x1, x2 = x[..., 0::2], x[..., 1::2] # channel pairs (0,1), (2,3), ...
out = torch.stack((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1)
return out.flatten(-2) # interleave the pairs back

The 64 channels of a head get grouped into 32 pairs. Each pair is a point on a 2D plane, and RoPE rotates that point by angle = position x frequency. Each pair has its own frequency. The first pair spins at frequency 1, one radian per position. The frequencies fall off geometrically toward 1/10000 across the later pairs. The fast pairs resolve nearby positions. The slow pairs barely move over thousands of tokens, which lets the head tell far positions apart.

The key property is what happens to the dot product. Rotate the query at position m by angle m x f, and the key at position n by angle n x f. The dot product of two rotated vectors depends only on the difference of the two angles. Rotate both by the same extra amount and the extra amount cancels. So the score between a query at m and a key at n depends only on m minus n: the relative distance.

The script checks this by planting the same q and k vectors at different positions:

RoPE score for 7 tokens apart at positions (10,3)=-3.7709, (110,103)=-3.7709, (507,500)=-3.7709
RoPE score for 1 token apart at (10,9) = -0.0418

Three different absolute positions, one distance of 7, one score. A distance of 1 gives a different score. The head can now learn “attend to whatever sits two tokens back” without ever being told an absolute position.

Three more details matter:

Now rerun the Holmes and Moriarty test with RoPE on:

RoPE: max difference in final '.' vector = 0.087948

The final ”.” now differs between the two sentences, by 0.087948. That number has no meaning yet, since the weights are still random. It only shows the two sentences are now distinguishable, and that is what the model needs before training can teach it who shot whom.

The Quadratic Bill

Back to the scores matrix. For T tokens, each head builds a T x T grid of scores. With 4 heads, that is 4 x T x T numbers. Double the sequence length and the matrix grows four times. This is the quadratic cost, and it is why a long context window costs so much.

scaling.py measures it for the shape Part 3’s model will use: 1 sequence, 4 heads of 64 dimensions. It compares two implementations. “Naive” builds the full score matrix, like attention.py does. “Fused” is F.scaled_dot_product_attention, which on a recent GPU picks a flash-attention kernel. First the CPU, an AMD Ryzen 7 255 with 8 threads in fp32:

CPU, 8 threads, torch 2.14.1+cpu, fp32
tokens score matrix naive ms fused ms speedup
256 1.0 MiB 0.67 0.17 3.9x
512 4.0 MiB 3.16 0.53 6.0x
1024 16.0 MiB 13.21 1.95 6.8x
2048 64.0 MiB 54.44 9.62 5.7x
4096 256.0 MiB 208.25 23.49 8.9x
8192 1024.0 MiB 759.44 92.05 8.3x

The score matrix column quadruples with each doubling, as predicted. The naive time does the same, roughly. The fused version is 4 to 9 times faster here. Across two runs on my machine the CPU speedup landed between about 4x and 9x, so treat the ratio as a range and the memory column as the fixed part.

Now the GPU, an RTX 3070 Laptop with 8 GB of VRAM in fp16, using torch 2.14.1+cu132 in the pytorch/pytorch:2.14.1-cuda13.2-cudnn9-runtime Docker image:

NVIDIA GeForce RTX 3070 Laptop GPU, torch 2.14.1+cu132, fp16
tokens score matrix naive ms fused ms speedup naive MiB fused MiB
256 0.5 MiB 0.17 0.04 4.5x 2 0.1
512 2.0 MiB 0.17 0.05 3.3x 6 1.3
1024 8.0 MiB 0.36 0.08 4.4x 25 0.5
2048 32.0 MiB 1.32 0.20 6.5x 100 1.0
4096 128.0 MiB 5.39 0.61 8.9x 400 2.1
8192 512.0 MiB 20.99 1.80 11.6x 1600 4.1
16384 2048.0 MiB 76.54 6.41 11.9x 6400 8.3
32768 8192.0 MiB OOM 22.14 - OOM 16.5

Four things stand out.

Naive memory explodes. At 8192 tokens, the naive path allocates 1600 MiB. The fused path uses 4.1 MiB. At 16384 tokens, naive takes 6400 MiB of an 8 GB card. The naive column runs about 3.1 times the score matrix because the code also holds the masked copy, the softmax output and a T x T boolean mask. At 32768 tokens the score matrix alone is 8192 MiB, 8 GiB, and the naive run dies with an out-of-memory error.

Fused memory grows linearly. Look at the last three fused entries: 4.1, 8.3, 16.5 MiB. Each doubling of the sequence doubles the memory. At 32768 tokens the fused kernel runs in 22.14 ms using 16.5 MiB. (The 512 row of that column is allocator noise, so ignore it.)

Time still grows about 3.5x per doubling. Fused attention at 8192, 16384 and 32768 takes 1.80, 6.41 and 22.14 ms, about 3.5x per doubling and closing in on the 4x that the arithmetic predicts. The arithmetic is still T x T. Fused attention saves memory traffic, not the quadratic work.

The GPU speedup reaches about 12x. From 1024 tokens up, the advantage grows with length (4.4x at 1024 tokens, 11.9x at 16384) because the naive path pays more and more to write and read that giant matrix.

How Flash Attention Avoids the Matrix

The FlashAttention paper describes the trick in plain terms. The GPU has a large pool of slow memory and a small pool of fast on-chip memory. The naive path writes the full T x T matrix out to slow memory, reads it back for the softmax, writes the result out again, then reads it for the weighted sum. Most of the time goes to that traffic.

Flash attention processes the scores in tiles that fit in on-chip memory. For each tile it computes the scores, updates a running softmax (a running maximum and a running sum, rescaled as new tiles arrive), and folds the tile into the output. The full matrix never exists. The result matches the naive answer within floating-point tolerance, and scaling.py checks that on a 64-token slice before it times anything.

What This Means for Inference

This table explains why long context costs VRAM. When a model generates text, it must keep the k and v vectors of every past token around, so each new token can attend to them. That store is the KV cache. It grows linearly with context length, and Part 4 builds one. For the practical side, see KV Cache Quantization, which shrinks that cache, and Context Window vs Token Limit, which covers what a context window actually limits. The score matrix is the other half of the bill during training and prompt processing, and fused kernels are why you do not pay it in memory.

What’s Next

Part 3 stacks this attention with an MLP into transformer blocks, adds the training loop, and trains a model under 30 million parameters, on the CPU and on the 8GB GPU. It also reruns Part 1’s Holmes, Watson and telegram cosine check on the trained embeddings, so you can see whether training changed which tokens sit close together.

The original design is in the Attention Is All You Need paper. The PyTorch function used throughout is documented under scaled_dot_product_attention.

Common Questions

Why divide attention scores by the square root of the head dimension?

Dividing by the square root of the head dimension keeps the variance of the dot products near 1. Without it, the scores grow with the dimension, softmax saturates toward a one-hot spike, and the gradients through it shrink. With 64-dimensional heads, the divisor is 8.

Does flash attention change the model’s output?

No. Flash attention computes the same math as the naive version, in a different order, so the results match within floating-point tolerance. The companion scripts scaling.py and attention.py assert that the fused and hand-rolled outputs agree. Only memory use and speed change.

How much VRAM does attention need for 32k tokens?

At 32,768 tokens, the naive score matrix alone takes 8192 MiB for 4 heads in fp16, so it overflows an 8 GB card. The fused kernel handled the same length using 16.5 MiB of extra memory in my test. Model weights and the KV cache still need their own VRAM on top.

Can I add RoPE to a model trained with learned position embeddings?

Not without more training. A model trained on learned position embeddings has weights that expect those added vectors, so swapping in rotations changes every attention score. The model needs fine-tuning or retraining to adapt. RoPE works as a design choice from the start of training, which is how Part 3 uses it.


Share this post on:

Send a Webmention

Written about this post on your own site? Send a webmention and it'll show up above once verified.


Previous Post
gVisor vs Firecracker vs Kata for Agents
Next Post
Network Namespaces by Hand, No Docker

Discussion

Powered by Garrul . Sign in with GitHub or Google, or post anonymously.

Related Posts