Attention Mechanisms and Transformers

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 nn, target position 2 to source position n−1n-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.

Build order, each step testable before the nextData, vocabulary and a toy reversal taskPositional encoding, checked by plotting itMulti-head attention, checked on shapesFeed-forward, encoder and decoder blocksMasks, training loop, greedy decoding
Overfit one batch to near-zero loss before touching the real data — it separates bugs from bad hyperparameters.

Step 0: the data

Thirteen tokens in the vocabulary: the digits 0–9 plus three special tokens.

Python
import torchimport torch.nn as nnimport torch.nn.functional as Fimport mathPAD, SOS, EOS = 10, 11, 12VOCAB = 13def make_batch(batch_size, min_len=4, max_len=10, device='cpu'):    """Returns (src, tgt). tgt is SOS + reversed(src) + EOS, all padded."""    lens = torch.randint(min_len, max_len + 1, (batch_size,))    T_src = int(lens.max())    T_tgt = T_src + 2                      # room for SOS and EOS    src = torch.full((batch_size, T_src), PAD, dtype=torch.long)    tgt = torch.full((batch_size, T_tgt), PAD, dtype=torch.long)    for i, L in enumerate(lens):        seq = torch.randint(0, 10, (L,))        src[i, :L] = seq        tgt[i, 0] = SOS        tgt[i, 1:L + 1] = seq.flip(0)      # the reversal        tgt[i, L + 1] = EOS    return src.to(device), tgt.to(device)src, tgt = make_batch(2, min_len=4, max_len=5)print(src[0])   # tensor([ 3,  1,  4,  1,  5])print(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

ln⁡13=2.565\ln 13 = 2.565

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\sigma^2/2 \approx 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:

PE(pos,2i)=sin⁡ ⁣(pos100002i/d),PE(pos,2i+1)=cos⁡ ⁣(pos100002i/d)\mathrm{PE}(pos, 2i) = \sin\!\left(\frac{pos}{10000^{2i/d}}\right), \qquad \mathrm{PE}(pos, 2i+1) = \cos\!\left(\frac{pos}{10000^{2i/d}}\right)

Python
class PositionalEncoding(nn.Module):    def __init__(self, d_model, max_len=512, dropout=0.1):        super().__init__()        self.dropout = nn.Dropout(dropout)        pe = torch.zeros(max_len, d_model)        pos = torch.arange(max_len).unsqueeze(1).float()        div = torch.exp(torch.arange(0, d_model, 2).float()                        * (-math.log(10000.0) / d_model))        pe[:, 0::2] = torch.sin(pos * div)        pe[:, 1::2] = torch.cos(pos * div)        self.register_buffer('pe', pe.unsqueeze(0))     # (1, max_len, d_model)    def forward(self, x):        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/d2i/d.

Step 2: multi-head attention

The core operation, in one formula:

Attention(Q,K,V)=softmax ⁣(QK⊤dk)V\mathrm{Attention}(Q,K,V) = \mathrm{softmax}\!\left(\frac{QK^{\top}}{\sqrt{d_k}}\right)V

Python
class MultiHeadAttention(nn.Module):    def __init__(self, d_model, num_heads, dropout=0.1):        super().__init__()        assert d_model % num_heads == 0        self.d_model, self.h = d_model, num_heads        self.d_k = d_model // num_heads        self.W_q = nn.Linear(d_model, d_model)        self.W_k = nn.Linear(d_model, d_model)        self.W_v = nn.Linear(d_model, d_model)        self.W_o = nn.Linear(d_model, d_model)        self.dropout = nn.Dropout(dropout)    def _split(self, x):        B, T, _ = x.shape        return x.view(B, T, self.h, self.d_k).transpose(1, 2)    # (B, h, T, d_k)    def forward(self, q, k, v, mask=None):        B, T_q, _ = q.shape        Q, K, V = self._split(self.W_q(q)), self._split(self.W_k(k)), self._split(self.W_v(v))        scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k)   # (B, h, T_q, T_k)        if mask is not None:            scores = scores.masked_fill(~mask, -1e9)        attn = self.dropout(torch.softmax(scores, dim=-1))        out = (attn @ V).transpose(1, 2).contiguous().view(B, T_q, self.d_model)        return self.W_o(out), attn

Three lines here are load-bearing, and each fails silently if you get it wrong.

The division by dk\sqrt{d_k}. If qq and kk have unit-variance components, q⋅kq\cdot k has variance dkd_k and standard deviation dk\sqrt{d_k}. With dk=16d_k = 16 that is 4, so logit gaps of 12 are ordinary. Softmax at a gap of 12 gives αmax⁡=e12/(1+e12)=0.9999939\alpha_{\max} = e^{12}/(1+e^{12}) = 0.9999939, and its gradient α(1−α)=6.1×10−6\alpha(1-\alpha) = 6.1\times10^{-6}. Divide by 4 and the gap becomes 3, giving α=0.953\alpha = 0.953 and gradient 0.0450.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 −∞-\infty gives 0/0=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:

Python
class FeedForward(nn.Module):    def __init__(self, d_model, d_ff, dropout=0.1):        super().__init__()        self.net = nn.Sequential(            nn.Linear(d_model, d_ff), nn.ReLU(),            nn.Dropout(dropout), nn.Linear(d_ff, d_model))    def forward(self, x):        return self.net(x)class EncoderLayer(nn.Module):    def __init__(self, d_model, h, d_ff, dropout=0.1):        super().__init__()        self.attn = MultiHeadAttention(d_model, h, dropout)        self.ff = FeedForward(d_model, d_ff, dropout)        self.n1, self.n2 = nn.LayerNorm(d_model), nn.LayerNorm(d_model)        self.drop = nn.Dropout(dropout)    def forward(self, x, mask=None):        a, _ = self.attn(x, x, x, mask)          # self-attention: q = k = v = x        x = self.n1(x + self.drop(a))            # residual, then normalise        x = self.n2(x + self.drop(self.ff(x)))        return xclass DecoderLayer(nn.Module):    def __init__(self, d_model, h, d_ff, dropout=0.1):        super().__init__()        self.self_attn  = MultiHeadAttention(d_model, h, dropout)        self.cross_attn = MultiHeadAttention(d_model, h, dropout)        self.ff = FeedForward(d_model, d_ff, dropout)        self.n1 = nn.LayerNorm(d_model)        self.n2 = nn.LayerNorm(d_model)        self.n3 = nn.LayerNorm(d_model)        self.drop = nn.Dropout(dropout)    def forward(self, x, memory, tgt_mask=None, src_mask=None):        a, _ = self.self_attn(x, x, x, tgt_mask)        x = self.n1(x + self.drop(a))        # Q from the decoder, K and V from the encoder output        c, cross = self.cross_attn(x, memory, memory, src_mask)        x = self.n2(x + self.drop(c))        x = self.n3(x + self.drop(self.ff(x)))        return x, cross

The residual connection x + sublayer(x) is what makes depth trainable. Its derivative is I+∂F/∂xI + \partial F/\partial 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.530.9^6 = 0.53 of the signal to layer 1, and at 24 layers only 0.924=0.080.9^{24} = 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

Python
def padding_mask(x, pad=PAD):    """(B, T) -> (B, 1, 1, T); True where the token is real."""    return (x != pad).unsqueeze(1).unsqueeze(2)def target_mask(x, pad=PAD):    """(B, T) -> (B, 1, T, T); padding AND causal."""    B, T = x.shape    pad_m = (x != pad).unsqueeze(1).unsqueeze(2)                       # (B,1,1,T)    causal = torch.tril(torch.ones(T, T, dtype=torch.bool,                                   device=x.device))[None, None]       # (1,1,T,T)    return pad_m & causalclass MiniTransformer(nn.Module):    def __init__(self, vocab=VOCAB, d_model=64, h=4, d_ff=256,                 num_layers=2, dropout=0.1, max_len=64):        super().__init__()        self.d_model = d_model        self.src_emb = nn.Embedding(vocab, d_model, padding_idx=PAD)        self.tgt_emb = nn.Embedding(vocab, d_model, padding_idx=PAD)        self.pos = PositionalEncoding(d_model, max_len, dropout)        self.enc = nn.ModuleList(EncoderLayer(d_model, h, d_ff, dropout)                                 for _ in range(num_layers))        self.dec = nn.ModuleList(DecoderLayer(d_model, h, d_ff, dropout)                                 for _ in range(num_layers))        self.generator = nn.Linear(d_model, vocab)        for p in self.parameters():            if p.dim() > 1:                nn.init.xavier_uniform_(p)    def encode(self, src, src_mask):        x = self.pos(self.src_emb(src) * math.sqrt(self.d_model))        for layer in self.enc:            x = layer(x, src_mask)        return x    def decode(self, tgt, memory, tgt_mask, src_mask):        x = self.pos(self.tgt_emb(tgt) * math.sqrt(self.d_model))        maps = []        for layer in self.dec:            x, cross = layer(x, memory, tgt_mask, src_mask)            maps.append(cross)        return x, maps    def forward(self, src, tgt):        src_mask = padding_mask(src)        tgt_mask = target_mask(tgt)        memory = self.encode(src, src_mask)        out, maps = self.decode(tgt, memory, tgt_mask, src_mask)        return self.generator(out), maps

The * 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)\sqrt{2/(V + d)}, which for V=13V=13, d=64d=64 is about 0.17. The positional encoding ranges over [−1,1][-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\sqrt{64} = 8 brings the token embeddings to a comparable scale.

Count the parameters — it is a useful habit, and the breakdown is instructive:

ComponentArithmeticParameters
Two embedding tables2×13×642 \times 13 \times 641,664
Attention block (each)4×(64×64+64)4 \times (64\times64 + 64)16,640
Feed-forward block (each)64 ⁣⋅ ⁣256+256+256 ⁣⋅ ⁣64+6464\!\cdot\!256 + 256 + 256\!\cdot\!64 + 6433,088
Encoder layers (×2\times 2)2×(16,640+33,088+256)2 \times (16{,}640 + 33{,}088 + 256)99,968
Decoder layers (×2\times 2)2×(2 ⁣⋅ ⁣16,640+33,088+384)2 \times (2\!\cdot\!16{,}640 + 33{,}088 + 384)133,504
Output projection64×13+1364 \times 13 + 13845
Total235,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 dffd_{ff} is the standard way to add capacity.

Step 5: the training loop

Python
def train(model, steps=3000, batch_size=64, lr=1e-3, device='cpu',          warmup=400, log_every=200):    model.to(device).train()    opt = torch.optim.Adam(model.parameters(), lr=lr, betas=(0.9, 0.98), eps=1e-9)    sched = torch.optim.lr_scheduler.LambdaLR(        opt, lambda s: min((s + 1) / warmup, 1.0))     # linear warmup, then flat    for step in range(1, steps + 1):        src, tgt = make_batch(batch_size, device=device)        tgt_in, tgt_out = tgt[:, :-1], tgt[:, 1:]      # SHIFT RIGHT        logits, _ = model(src, tgt_in)        loss = F.cross_entropy(            logits.reshape(-1, VOCAB), tgt_out.reshape(-1),            ignore_index=PAD)        opt.zero_grad(set_to_none=True)        loss.backward()        gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)        opt.step()        sched.step()        if step % log_every == 0 or step == 1:            print(f"step {step:5d}  loss {loss.item():.4f}  "                  f"grad_norm {gnorm:.2f}  lr {sched.get_last_lr()[0]:.2e}")    return model

The two lines tgt_in = tgt[:, :-1] and tgt_out = tgt[:, 1:] are the most important in the file. Laid out for the example above:

Position012345
Decoder input<sos>51413
Label to predict51413<eos>

Each position sees the tokens up to and including itself and must predict the next one. Feed the unshifted target and position ii sees token ii while being asked to produce token ii: 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:

Text
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-03

Step 1 at 3.35 is ln⁡13=2.565\ln 13 = 2.565 plus the σ2/2≈0.8\sigma^2/2 \approx 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

Python
@torch.no_grad()def greedy_decode(model, src, max_len=20):    model.eval()    src_mask = padding_mask(src)    memory = model.encode(src, src_mask)               # encoder runs ONCE    ys = torch.full((src.size(0), 1), SOS, dtype=torch.long, device=src.device)    for _ in range(max_len - 1):        out, maps = model.decode(ys, memory, target_mask(ys), src_mask)        next_tok = model.generator(out[:, -1]).argmax(-1, keepdim=True)        ys = torch.cat([ys, next_tok], dim=1)        if (next_tok == EOS).all():            break    return ys, mapssrc, tgt = make_batch(1, min_len=6, max_len=6)pred, maps = greedy_decode(model, src)print("source    ", src[0].tolist())          # [7, 2, 9, 0, 4, 6]print("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-TT output takes TT decoder forward passes, each attending over up to TT positions, so decoding is O(T2)O(T^2) 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.

Python
import matplotlib.pyplot as pltdef show_cross_attention(maps, src, pred, layer=-1, head=0):    # maps[layer]: (B, h, T_tgt, T_src)    w = maps[layer][0, head].detach().cpu().numpy()    fig, ax = plt.subplots(figsize=(6, 5))    im = ax.imshow(w, cmap='viridis', aspect='auto', vmin=0, vmax=1)    ax.set_xticks(range(src.size(1)), [str(t) for t in src[0].tolist()])    # row i is the query at input position i, which produced token i+1    ax.set_yticks(range(w.shape[0]), [str(t) for t in pred[0, 1:w.shape[0] + 1].tolist()])    ax.set_xlabel('source position')    ax.set_ylabel('generated token')    fig.colorbar(im)    plt.tight_layout()    return fig

For source 7 2 9 0 4 6 a trained model produces a cross-attention matrix close to:

Text
              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.01

Every row sums to 1.00 — check the first: 0.01+0.01+0.02+0.03+0.05+0.88=1.000.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.

SymptomCauseFix
Loss under 0.05 within 200 steps; generation emits one token foreverTarget not shifted right — the model is copying its inputtgt[:, :-1] in, tgt[:, 1:] as labels
Training loss excellent, generation is garbageCausal mask inverted (triu instead of tril) or missingPrint the mask; run the leakage test below
Loss falls to about 1.2 and stopsMissing dk\sqrt{d_k} — softmax saturated from step onePrint max⁡\max attention weight per row; above 0.999 at init means no scaling
Loss stuck near ln⁡V=2.565\ln V = 2.565 and never movesLearning rate far too low, or gradients not reaching the modelPrint grad_norm; if it is 0, an input is detached or wrapped in no_grad
nan after a few hundred stepsAn entirely masked row producing 0/00/0 in softmax, or fp16 overflowUse -1e9 not -inf; clip gradients to 1.0
view() raises "not contiguous"Reshaping straight after transposeInsert .contiguous(), or use .reshape()
Model ignores the source; output is fluent but unrelatedCross-attention arguments swapped, or memory not passed to every layerConfirm cross_attn(x, memory, memory); plot the map
Works at length 6, fails at length 30max_len in the positional encoding is too small, or masks built for one length are being reusedBuild masks from the actual tensor shapes each step
Loss oscillates violently from step 1No warmupLinear warmup over a few hundred steps

Two tests catch most of the table:

Python
@torch.no_grad()def test_no_future_leak(model):    model.eval()    src, tgt = make_batch(1, min_len=8, max_len=8)    tgt_in = tgt[:, :-1]    a, _ = model(src, tgt_in)    t2 = tgt_in.clone()    t2[0, -1] = (t2[0, -1] + 1) % 10           # change only the LAST input token    b, _ = model(src, t2)    assert torch.allclose(a[0, :-1], b[0, :-1], atol=1e-6), "future leakage"    print("causal masking OK")def test_overfit_one_batch(model, steps=400):    """A correct model must drive ONE fixed batch to near-zero loss.    Pass a fresh MiniTransformer(dropout=0.0): with dropout on, it cannot memorise."""    model.train()    src, tgt = make_batch(8)    opt = torch.optim.Adam(model.parameters(), lr=1e-3)    for _ in range(steps):        logits, _ = model(src, tgt[:, :-1])        loss = F.cross_entropy(logits.reshape(-1, VOCAB),                               tgt[:, 1:].reshape(-1), ignore_index=PAD)        opt.zero_grad(set_to_none=True); loss.backward(); opt.step()    print(f"overfit loss {loss.item():.4f}")   # should be well under 0.01    assert loss.item() < 0.05

Run 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−510^{-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:

StageWhat to verify before moving on
Data generationPrint five examples and read them. The reversal, the special tokens, the padding.
ShapesOne forward pass with batch 2. Every intermediate shape is what you expect.
Initial lossClose to ln⁡V\ln V — here 2.6 to 3.4. Far above or far below means a bug.
MaskingThe leakage test passes.
OptimisationThe overfit-one-batch test reaches loss under 0.01.
GeneralisationHeld-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.