Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Implementing a Mini Transformer from Scratch
You can read the transformer paper five times and still not know whether you understand it. The test is different: can you write one that trains, from an empty file, and can you tell what is wrong when it does not?
So we will build one. Not a toy that only prints shapes — a complete encoder-decoder that learns a real mapping, generates output autoregressively, and lets you look at its attention. The task is deliberately small enough that it trains on a laptop CPU in under two minutes, which means you can break it, fix it, and see the effect immediately.
The task: reverse a sequence of digits. Input 3 1 4 1 5, output 5 1 4 1 3. This is a better choice than copying, because the correct cross-attention pattern is an anti-diagonal — target position 1 must attend to source position n, target position 2 to source position n−1, and so on. A model that has merely learned to copy positionally will show a diagonal and be visibly wrong. The attention map becomes a correctness check you can see.
Step 0: the data
Thirteen tokens in the vocabulary: the digits 0–9 plus three special tokens.
1import torch2import torch.nn as nn3import torch.nn.functional as F4import math56PAD, SOS, EOS = 10, 11, 127VOCAB = 1389def make_batch(batch_size, min_len=4, max_len=10, device='cpu'):10 """Returns (src, tgt). tgt is SOS + reversed(src) + EOS, all padded."""11 lens = torch.randint(min_len, max_len + 1, (batch_size,))12 T_src = int(lens.max())13 T_tgt = T_src + 2 # room for SOS and EOS1415 src = torch.full((batch_size, T_src), PAD, dtype=torch.long)16 tgt = torch.full((batch_size, T_tgt), PAD, dtype=torch.long)1718 for i, L in enumerate(lens):19 seq = torch.randint(0, 10, (L,))20 src[i, :L] = seq21 tgt[i, 0] = SOS22 tgt[i, 1:L + 1] = seq.flip(0) # the reversal23 tgt[i, L + 1] = EOS2425 return src.to(device), tgt.to(device)2627src, tgt = make_batch(2, min_len=4, max_len=5)28print(src[0]) # tensor([ 3, 1, 4, 1, 5])29print(tgt[0]) # tensor([11, 5, 1, 4, 1, 3, 12])Sanity-check that by hand: source 3 1 4 1 5, target <sos> 5 1 4 1 3 <eos>. The reversal is right and the special tokens bracket it.
A useful number before we start: a model that has learned nothing assigns uniform probability over 13 tokens, so cross-entropy should begin near
In practice this model starts a little higher, around 3.3. The randomly initialised output layer gives logits with a standard deviation of about 1.3 rather than zero, and random logits add roughly σ2/2≈0.8 to the expected loss. Anything from about 2.5 to 3.5 is healthy. If your first loss is 8.0, the output layer or the initialisation is wrong. If it is 0.4, labels are leaking into the input. This single number catches a surprising number of bugs before you have burned an hour.
Step 1: positional encoding
Attention has no notion of order — permute the input rows and the output rows permute identically, because nothing in the attention formula references an index. Position must therefore be added to the embeddings before the first layer. We use the fixed sinusoidal formula:
1class PositionalEncoding(nn.Module):2 def __init__(self, d_model, max_len=512, dropout=0.1):3 super().__init__()4 self.dropout = nn.Dropout(dropout)5 pe = torch.zeros(max_len, d_model)6 pos = torch.arange(max_len).unsqueeze(1).float()7 div = torch.exp(torch.arange(0, d_model, 2).float()8 * (-math.log(10000.0) / d_model))9 pe[:, 0::2] = torch.sin(pos * div)10 pe[:, 1::2] = torch.cos(pos * div)11 self.register_buffer('pe', pe.unsqueeze(0)) # (1, max_len, d_model)1213 def forward(self, x):14 return self.dropout(x + self.pe[:, :x.size(1)])register_buffer rather than nn.Parameter: this table is a constant. It should be saved in the state dict and moved by .to(device), but it must never receive a gradient.
The div term is computed in log space, as in the reference implementations. Writing 10000 ** (-2*i/d) directly gives the same values; what you must not get wrong is the exponent 2i/d.
Step 2: multi-head attention
The core operation, in one formula:
1class MultiHeadAttention(nn.Module):2 def __init__(self, d_model, num_heads, dropout=0.1):3 super().__init__()4 assert d_model % num_heads == 05 self.d_model, self.h = d_model, num_heads6 self.d_k = d_model // num_heads7 self.W_q = nn.Linear(d_model, d_model)8 self.W_k = nn.Linear(d_model, d_model)9 self.W_v = nn.Linear(d_model, d_model)10 self.W_o = nn.Linear(d_model, d_model)11 self.dropout = nn.Dropout(dropout)1213 def _split(self, x):14 B, T, _ = x.shape15 return x.view(B, T, self.h, self.d_k).transpose(1, 2) # (B, h, T, d_k)1617 def forward(self, q, k, v, mask=None):18 B, T_q, _ = q.shape19 Q, K, V = self._split(self.W_q(q)), self._split(self.W_k(k)), self._split(self.W_v(v))2021 scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k) # (B, h, T_q, T_k)22 if mask is not None:23 scores = scores.masked_fill(~mask, -1e9)2425 attn = self.dropout(torch.softmax(scores, dim=-1))26 out = (attn @ V).transpose(1, 2).contiguous().view(B, T_q, self.d_model)27 return self.W_o(out), attnThree lines here are load-bearing, and each fails silently if you get it wrong.
The division by dk. If q and k have unit-variance components, q⋅k has variance dk and standard deviation dk. With dk=16 that is 4, so logit gaps of 12 are ordinary. Softmax at a gap of 12 gives αmax=e12/(1+e12)=0.9999939, and its gradient α(1−α)=6.1×10−6. Divide by 4 and the gap becomes 3, giving α=0.953 and gradient 0.045 — nearly four orders of magnitude healthier. Omit the division and the model plateaus almost immediately.
dim=-1 in the softmax. Each row of the score matrix is one query against all keys, so normalisation runs over keys, the last axis. Using dim=-2 produces identical shapes, no error, and a meaningless model.
-1e9 rather than float('-inf'). If an entire row is masked — which happens for fully-padded positions — softmax over all −∞ gives 0/0= nan, and one nan poisons every parameter on the backward pass. A large finite negative number gives a uniform row instead: useless, but harmless.
Step 3: feed-forward, encoder, decoder
The feed-forward network is applied independently to every position — same weights, no mixing between positions. It expands by a factor of four and comes back:
1class FeedForward(nn.Module):2 def __init__(self, d_model, d_ff, dropout=0.1):3 super().__init__()4 self.net = nn.Sequential(5 nn.Linear(d_model, d_ff), nn.ReLU(),6 nn.Dropout(dropout), nn.Linear(d_ff, d_model))78 def forward(self, x):9 return self.net(x)1011class EncoderLayer(nn.Module):12 def __init__(self, d_model, h, d_ff, dropout=0.1):13 super().__init__()14 self.attn = MultiHeadAttention(d_model, h, dropout)15 self.ff = FeedForward(d_model, d_ff, dropout)16 self.n1, self.n2 = nn.LayerNorm(d_model), nn.LayerNorm(d_model)17 self.drop = nn.Dropout(dropout)1819 def forward(self, x, mask=None):20 a, _ = self.attn(x, x, x, mask) # self-attention: q = k = v = x21 x = self.n1(x + self.drop(a)) # residual, then normalise22 x = self.n2(x + self.drop(self.ff(x)))23 return x2425class DecoderLayer(nn.Module):26 def __init__(self, d_model, h, d_ff, dropout=0.1):27 super().__init__()28 self.self_attn = MultiHeadAttention(d_model, h, dropout)29 self.cross_attn = MultiHeadAttention(d_model, h, dropout)30 self.ff = FeedForward(d_model, d_ff, dropout)31 self.n1 = nn.LayerNorm(d_model)32 self.n2 = nn.LayerNorm(d_model)33 self.n3 = nn.LayerNorm(d_model)34 self.drop = nn.Dropout(dropout)3536 def forward(self, x, memory, tgt_mask=None, src_mask=None):37 a, _ = self.self_attn(x, x, x, tgt_mask)38 x = self.n1(x + self.drop(a))3940 # Q from the decoder, K and V from the encoder output41 c, cross = self.cross_attn(x, memory, memory, src_mask)42 x = self.n2(x + self.drop(c))4344 x = self.n3(x + self.drop(self.ff(x)))45 return x, crossThe residual connection x + sublayer(x) is what makes depth trainable. Its derivative is I+∂F/∂x; the identity term means the gradient reaching early layers does not decay geometrically with depth. Without it, a 6-layer stack with an average per-layer gradient factor of 0.9 delivers only 0.96=0.53 of the signal to layer 1, and at 24 layers only 0.924=0.08.
The line self.cross_attn(x, memory, memory, src_mask) is the entire encoder-decoder connection. Query from the decoder, key and value from the encoder. Reverse the arguments and, when source and target lengths happen to match, you get no error at all — just a model that learns the wrong thing.
Step 4: masks and the full model
1def padding_mask(x, pad=PAD):2 """(B, T) -> (B, 1, 1, T); True where the token is real."""3 return (x != pad).unsqueeze(1).unsqueeze(2)45def target_mask(x, pad=PAD):6 """(B, T) -> (B, 1, T, T); padding AND causal."""7 B, T = x.shape8 pad_m = (x != pad).unsqueeze(1).unsqueeze(2) # (B,1,1,T)9 causal = torch.tril(torch.ones(T, T, dtype=torch.bool,10 device=x.device))[None, None] # (1,1,T,T)11 return pad_m & causal1213class MiniTransformer(nn.Module):14 def __init__(self, vocab=VOCAB, d_model=64, h=4, d_ff=256,15 num_layers=2, dropout=0.1, max_len=64):16 super().__init__()17 self.d_model = d_model18 self.src_emb = nn.Embedding(vocab, d_model, padding_idx=PAD)19 self.tgt_emb = nn.Embedding(vocab, d_model, padding_idx=PAD)20 self.pos = PositionalEncoding(d_model, max_len, dropout)2122 self.enc = nn.ModuleList(EncoderLayer(d_model, h, d_ff, dropout)23 for _ in range(num_layers))24 self.dec = nn.ModuleList(DecoderLayer(d_model, h, d_ff, dropout)25 for _ in range(num_layers))26 self.generator = nn.Linear(d_model, vocab)2728 for p in self.parameters():29 if p.dim() > 1:30 nn.init.xavier_uniform_(p)3132 def encode(self, src, src_mask):33 x = self.pos(self.src_emb(src) * math.sqrt(self.d_model))34 for layer in self.enc:35 x = layer(x, src_mask)36 return x3738 def decode(self, tgt, memory, tgt_mask, src_mask):39 x = self.pos(self.tgt_emb(tgt) * math.sqrt(self.d_model))40 maps = []41 for layer in self.dec:42 x, cross = layer(x, memory, tgt_mask, src_mask)43 maps.append(cross)44 return x, maps4546 def forward(self, src, tgt):47 src_mask = padding_mask(src)48 tgt_mask = target_mask(tgt)49 memory = self.encode(src, src_mask)50 out, maps = self.decode(tgt, memory, tgt_mask, src_mask)51 return self.generator(out), mapsThe * math.sqrt(self.d_model) on the embeddings is easy to dismiss as superstition. It is not. Xavier initialisation gives embedding entries a standard deviation around 2/(V+d), which for V=13, d=64 is about 0.17. The positional encoding ranges over [−1,1] in every coordinate. Without the scaling the position signal is roughly four times larger than the token signal, and early training is dominated by position. Multiplying by 64=8 brings the token embeddings to a comparable scale.
Count the parameters — it is a useful habit, and the breakdown is instructive:
| Component | Arithmetic | Parameters |
|---|---|---|
| Two embedding tables | 2×13×64 | 1,664 |
| Attention block (each) | 4×(64×64+64) | 16,640 |
| Feed-forward block (each) | 64⋅256+256+256⋅64+64 | 33,088 |
| Encoder layers (×2) | 2×(16,640+33,088+256) | 99,968 |
| Decoder layers (×2) | 2×(2⋅16,640+33,088+384) | 133,504 |
| Output projection | 64×13+13 | 845 |
| Total | 235,981 |
Note the ratio inside a layer: 33,088 feed-forward parameters against 16,640 for attention. Two thirds of a transformer's weights sit in the feed-forward networks, not in attention. That surprises people, and it is why widening dff is the standard way to add capacity.
Step 5: the training loop
1def train(model, steps=3000, batch_size=64, lr=1e-3, device='cpu',2 warmup=400, log_every=200):3 model.to(device).train()4 opt = torch.optim.Adam(model.parameters(), lr=lr, betas=(0.9, 0.98), eps=1e-9)5 sched = torch.optim.lr_scheduler.LambdaLR(6 opt, lambda s: min((s + 1) / warmup, 1.0)) # linear warmup, then flat78 for step in range(1, steps + 1):9 src, tgt = make_batch(batch_size, device=device)10 tgt_in, tgt_out = tgt[:, :-1], tgt[:, 1:] # SHIFT RIGHT1112 logits, _ = model(src, tgt_in)13 loss = F.cross_entropy(14 logits.reshape(-1, VOCAB), tgt_out.reshape(-1),15 ignore_index=PAD)1617 opt.zero_grad(set_to_none=True)18 loss.backward()19 gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)20 opt.step()21 sched.step()2223 if step % log_every == 0 or step == 1:24 print(f"step {step:5d} loss {loss.item():.4f} "25 f"grad_norm {gnorm:.2f} lr {sched.get_last_lr()[0]:.2e}")26 return modelThe two lines tgt_in = tgt[:, :-1] and tgt_out = tgt[:, 1:] are the most important in the file. Laid out for the example above:
| Position | 0 | 1 | 2 | 3 | 4 | 5 |
|---|---|---|---|---|---|---|
| Decoder input | <sos> | 5 | 1 | 4 | 1 | 3 |
| Label to predict | 5 | 1 | 4 | 1 | 3 | <eos> |
Each position sees the tokens up to and including itself and must predict the next one. Feed the unshifted target and position i sees token i while being asked to produce token i: the model learns the identity function, the loss drops below 0.05 within a couple of hundred steps, and generation emits <sos> forever.
A typical healthy run on this task:
step 1 loss 3.3463 grad_norm 5.53 lr 5.00e-06step 200 loss 1.5749 grad_norm 1.18 lr 5.02e-04step 400 loss 0.9043 grad_norm 1.49 lr 1.00e-03step 1000 loss 0.2748 grad_norm 2.33 lr 1.00e-03step 2000 loss 0.0882 grad_norm 1.40 lr 1.00e-03step 3000 loss 0.0412 grad_norm 1.18 lr 1.00e-03Step 1 at 3.35 is ln13=2.565 plus the σ2/2≈0.8 offset from the random output layer, as predicted above. That agreement tells you the output layer, the vocabulary size and the loss are all consistent before a single useful gradient has been taken. The later losses bounce around rather than falling smoothly because every step draws a fresh random batch and dropout is on; what matters is the trend. After this run, greedy decoding reverses 100% of 500 fresh random sequences.
Step 6: greedy decoding
1@torch.no_grad()2def greedy_decode(model, src, max_len=20):3 model.eval()4 src_mask = padding_mask(src)5 memory = model.encode(src, src_mask) # encoder runs ONCE67 ys = torch.full((src.size(0), 1), SOS, dtype=torch.long, device=src.device)8 for _ in range(max_len - 1):9 out, maps = model.decode(ys, memory, target_mask(ys), src_mask)10 next_tok = model.generator(out[:, -1]).argmax(-1, keepdim=True)11 ys = torch.cat([ys, next_tok], dim=1)12 if (next_tok == EOS).all():13 break14 return ys, maps1516src, tgt = make_batch(1, min_len=6, max_len=6)17pred, maps = greedy_decode(model, src)18print("source ", src[0].tolist()) # [7, 2, 9, 0, 4, 6]19print("prediction", pred[0].tolist()) # [11, 6, 4, 0, 9, 2, 7, 12]Two things distinguish this from training. The encoder runs once, outside the loop — recomputing the memory at every decode step is a common and entirely wasteful mistake. And only the last position's logits matter: the earlier positions are recomputed identically every step, which is the redundancy a KV cache exists to remove.
An honest accounting of cost: generating a length-T output takes T decoder forward passes, each attending over up to T positions, so decoding is O(T2) work even though training the same sequence was one pass. This is why generation is slow and why so much inference engineering is about caching.
Step 7: looking at cross-attention
This is where the reversal task pays off.
1import matplotlib.pyplot as plt23def show_cross_attention(maps, src, pred, layer=-1, head=0):4 # maps[layer]: (B, h, T_tgt, T_src)5 w = maps[layer][0, head].detach().cpu().numpy()6 fig, ax = plt.subplots(figsize=(6, 5))7 im = ax.imshow(w, cmap='viridis', aspect='auto', vmin=0, vmax=1)8 ax.set_xticks(range(src.size(1)), [str(t) for t in src[0].tolist()])9 # row i is the query at input position i, which produced token i+110 ax.set_yticks(range(w.shape[0]), [str(t) for t in pred[0, 1:w.shape[0] + 1].tolist()])11 ax.set_xlabel('source position')12 ax.set_ylabel('generated token')13 fig.colorbar(im)14 plt.tight_layout()15 return figFor source 7 2 9 0 4 6 a trained model produces a cross-attention matrix close to:
src: 7 2 9 0 4 6generated 6 0.01 0.01 0.02 0.03 0.05 0.88generated 4 0.01 0.02 0.03 0.06 0.85 0.03generated 0 0.02 0.03 0.05 0.84 0.04 0.02generated 9 0.03 0.06 0.83 0.04 0.02 0.02generated 2 0.05 0.86 0.04 0.02 0.02 0.01generated 7 0.87 0.05 0.03 0.02 0.02 0.01Every row sums to 1.00 — check the first: 0.01+0.01+0.02+0.03+0.05+0.88=1.00. And the bright band runs from top-right to bottom-left: a clean anti-diagonal. The model has discovered the reversal alignment on its own, from nothing but input-output pairs.
If you instead see a bright main diagonal, the model is copying rather than reversing — which means your data generation forgot the .flip(0). If you see a uniform grey field, cross-attention is not being used at all: check that memory is reaching every decoder layer.
Cross-attention maps are the cheapest correctness test you have on a sequence-to-sequence model, because the correct alignment is usually something you can recognise by eye.
Debugging: symptom to cause
These are the failures that actually happen, ordered by how often they do.
| Symptom | Cause | Fix |
|---|---|---|
| Loss under 0.05 within 200 steps; generation emits one token forever | Target not shifted right — the model is copying its input | tgt[:, :-1] in, tgt[:, 1:] as labels |
| Training loss excellent, generation is garbage | Causal mask inverted (triu instead of tril) or missing | Print the mask; run the leakage test below |
| Loss falls to about 1.2 and stops | Missing dk — softmax saturated from step one | Print max attention weight per row; above 0.999 at init means no scaling |
| Loss stuck near lnV=2.565 and never moves | Learning rate far too low, or gradients not reaching the model | Print grad_norm; if it is 0, an input is detached or wrapped in no_grad |
nan after a few hundred steps | An entirely masked row producing 0/0 in softmax, or fp16 overflow | Use -1e9 not -inf; clip gradients to 1.0 |
view() raises "not contiguous" | Reshaping straight after transpose | Insert .contiguous(), or use .reshape() |
| Model ignores the source; output is fluent but unrelated | Cross-attention arguments swapped, or memory not passed to every layer | Confirm cross_attn(x, memory, memory); plot the map |
| Works at length 6, fails at length 30 | max_len in the positional encoding is too small, or masks built for one length are being reused | Build masks from the actual tensor shapes each step |
| Loss oscillates violently from step 1 | No warmup | Linear warmup over a few hundred steps |
Two tests catch most of the table:
1@torch.no_grad()2def test_no_future_leak(model):3 model.eval()4 src, tgt = make_batch(1, min_len=8, max_len=8)5 tgt_in = tgt[:, :-1]6 a, _ = model(src, tgt_in)78 t2 = tgt_in.clone()9 t2[0, -1] = (t2[0, -1] + 1) % 10 # change only the LAST input token10 b, _ = model(src, t2)1112 assert torch.allclose(a[0, :-1], b[0, :-1], atol=1e-6), "future leakage"13 print("causal masking OK")1415def test_overfit_one_batch(model, steps=400):16 """A correct model must drive ONE fixed batch to near-zero loss.17 Pass a fresh MiniTransformer(dropout=0.0): with dropout on, it cannot memorise."""18 model.train()19 src, tgt = make_batch(8)20 opt = torch.optim.Adam(model.parameters(), lr=1e-3)21 for _ in range(steps):22 logits, _ = model(src, tgt[:, :-1])23 loss = F.cross_entropy(logits.reshape(-1, VOCAB),24 tgt[:, 1:].reshape(-1), ignore_index=PAD)25 opt.zero_grad(set_to_none=True); loss.backward(); opt.step()26 print(f"overfit loss {loss.item():.4f}") # should be well under 0.0127 assert loss.item() < 0.05Run the overfit test on a fresh model built with dropout=0.0; with dropout left at 0.1 the loss stalls around 0.07 and the assertion fails for the wrong reason. With dropout off it reaches about 0.0005 in 400 steps. The overfit test is the single highest-value thing in this lesson. A model that cannot memorise eight examples has a wiring bug, full stop — no learning rate, no dataset size and no architecture change will rescue it. Two minutes here saves days of staring at a plateau.
What this means when you build something
In production you would not write most of this. nn.TransformerEncoderLayer (pass norm_first=True for pre-norm), nn.MultiheadAttention and F.scaled_dot_product_attention are faster, better tested, and handle the numerics you would otherwise get subtly wrong. Use them.
But keep the hand-written version in your test suite and assert the two agree to within 10−5 on a fixed input. The reason is specific: the fused kernels never materialise the attention matrix, so when something is wrong you have no intermediates to inspect. The slow explicit path is the one whose numbers you can print.
Build in this order, and check something at each stage rather than at the end:
| Stage | What to verify before moving on |
|---|---|
| Data generation | Print five examples and read them. The reversal, the special tokens, the padding. |
| Shapes | One forward pass with batch 2. Every intermediate shape is what you expect. |
| Initial loss | Close to lnV — here 2.6 to 3.4. Far above or far below means a bug. |
| Masking | The leakage test passes. |
| Optimisation | The overfit-one-batch test reaches loss under 0.01. |
| Generalisation | Held-out accuracy, and a cross-attention map that shows the alignment you expect. |
Once this works, scale it by changing four numbers — d_model, num_layers, d_ff, num_heads — and swapping the data. The core mechanism does not change between this 236,000-parameter model and a 7-billion-parameter one. What does change are component choices: current large language models are decoder-only and use pre-norm (usually with RMSNorm, a cheaper LayerNorm), rotary position embeddings instead of sinusoidal ones, a gated feed-forward layer such as SwiGLU, and grouped-query attention running on a fused kernel. Each is a swap of one part, not a different idea. The rest is arithmetic, hardware, and data.