{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Lesson 46: The Transformer Architecture\n", "\n", "Lesson 45 built one attention operation. A **Transformer** (Vaswani et al., 2017★) wraps that operation into a reusable block and stacks many of them. This lesson fills in the three pieces attention alone is missing: **multiple attention heads** (so a layer can track several kinds of relationships at once), **positional encoding** (so order isn't invisible to a mechanism that otherwise treats a sequence as an unordered set), and the **encoder block** (attention plus a per-position feedforward network, wired together with residual connections and normalization)." ] }, { "cell_type": "code", "execution_count": null, "id": "ee266484", "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import torch\n", "import torch.nn as nn\n", "import torch.nn.functional as F\n", "import matplotlib.pyplot as plt" ] }, { "cell_type": "markdown", "id": "beb43435", "metadata": {}, "source": [ "## Multi-head attention\n", "\n", "A single attention operation computes one similarity pattern between every pair of positions. **Multi-head attention** splits the model dimension into several smaller chunks (heads), runs attention independently within each, and concatenates the results — letting different heads specialize in different kinds of relationships (e.g. one head tracking nearby positions, another tracking a specific long-range dependency) instead of averaging everything into one pattern." ] }, { "cell_type": "code", "execution_count": null, "id": "2bd3929c", "metadata": {}, "outputs": [], "source": [ "D, H = 16, 4 # model dimension, number of heads\n", "mha_torch = nn.MultiheadAttention(D, H, batch_first=True)\n", "\n", "def multihead_attention(x, mha):\n", " B, T, _ = x.shape\n", " d_head = D // H\n", " Wq, Wk, Wv = mha.in_proj_weight.chunk(3, dim=0)\n", " bq, bk, bv = mha.in_proj_bias.chunk(3, dim=0)\n", " Q = (x @ Wq.T + bq).view(B, T, H, d_head).transpose(1, 2)\n", " K = (x @ Wk.T + bk).view(B, T, H, d_head).transpose(1, 2)\n", " V = (x @ Wv.T + bv).view(B, T, H, d_head).transpose(1, 2)\n", " scores = Q @ K.transpose(-2, -1) / np.sqrt(d_head)\n", " weights = F.softmax(scores, dim=-1)\n", " out = (weights @ V).transpose(1, 2).reshape(B, T, D)\n", " return out @ mha.out_proj.weight.T + mha.out_proj.bias\n", "\n", "torch.manual_seed(0)\n", "x = torch.randn(2, 5, D)\n", "out_manual = multihead_attention(x, mha_torch)\n", "out_torch, _ = mha_torch(x, x, x, need_weights=False)\n", "\n", "print(f'model dim = {D}, heads = {H}, dim per head = {D // H}')\n", "print(f'max abs diff vs nn.MultiheadAttention: {(out_manual - out_torch).abs().max().item():.2e}')" ] }, { "cell_type": "markdown", "id": "f27c8a37", "metadata": {}, "source": [ "## Positional encoding\n", "\n", "Attention computes similarity between content vectors — nothing about the formula from Lesson 45 refers to *where* in the sequence a position sits. Permute the input sequence, and attention's output permutes along with it, identically: it is fundamentally a set operation, blind to order. **Positional encoding** fixes this by adding a unique, deterministic pattern to each position before attention runs, so position becomes part of what gets compared. The original Transformer paper's choice is a fixed (not learned) sinusoid at multiple frequencies:\n", "\n", "$$PE_{(pos, 2i)} = \\sin(pos / 10000^{2i/D}), \\qquad PE_{(pos, 2i+1)} = \\cos(pos / 10000^{2i/D})$$" ] }, { "cell_type": "code", "execution_count": null, "id": "5ce2a493", "metadata": {}, "outputs": [], "source": [ "def positional_encoding(T, D):\n", " pos = torch.arange(T).unsqueeze(1).float()\n", " i = torch.arange(D).unsqueeze(0).float()\n", " angle_rates = 1.0 / (10000 ** (2 * (i // 2) / D))\n", " angles = pos * angle_rates\n", " pe = torch.zeros(T, D)\n", " pe[:, 0::2] = torch.sin(angles[:, 0::2])\n", " pe[:, 1::2] = torch.cos(angles[:, 1::2])\n", " return pe\n", "\n", "pe = positional_encoding(50, 32)\n", "plt.figure(figsize=(6, 4))\n", "plt.imshow(pe.numpy().T, cmap='RdBu', aspect='auto')\n", "plt.xlabel('position'); plt.ylabel('encoding dimension')\n", "plt.title('Sinusoidal positional encoding')\n", "plt.colorbar(fraction=0.046)\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "9373cea9", "metadata": {}, "source": [ "### Does this actually matter? A task that requires knowing order\n", "\n", "Build a task that's impossible to solve from content alone: a fixed vector `A` and a fixed vector `B` appear at two random positions in a noisy sequence, and the label is simply \"does `A` appear before `B`?\" A model that mean-pools attention output over the sequence can only tell *which two tokens are present*, never their order, unless position is somehow injected." ] }, { "cell_type": "code", "execution_count": null, "id": "2df779e4", "metadata": {}, "outputs": [], "source": [ "T = 8\n", "vec_a, vec_b = torch.randn(D), torch.randn(D)\n", "\n", "def make_order_dataset(rng, n):\n", " X, y = [], []\n", " for _ in range(n):\n", " pos_a, pos_b = rng.choice(T, size=2, replace=False)\n", " seq = torch.randn(T, D) * 0.1\n", " seq[pos_a] = vec_a\n", " seq[pos_b] = vec_b\n", " X.append(seq)\n", " y.append(1.0 if pos_a < pos_b else 0.0)\n", " return torch.stack(X), torch.tensor(y, dtype=torch.float32)\n", "\n", "rng = np.random.default_rng(3)\n", "X_train, y_train = make_order_dataset(rng, 400)\n", "X_test, y_test = make_order_dataset(rng, 150)\n", "\n", "class TinyAttnClassifier(nn.Module):\n", " def __init__(self, use_pos_enc):\n", " super().__init__()\n", " self.use_pos_enc = use_pos_enc\n", " self.mha = nn.MultiheadAttention(D, 4, batch_first=True)\n", " self.fc = nn.Sequential(nn.Linear(D, 16), nn.ReLU(), nn.Linear(16, 1))\n", " if use_pos_enc:\n", " self.register_buffer('pe', positional_encoding(T, D))\n", "\n", " def forward(self, x):\n", " if self.use_pos_enc:\n", " x = x + self.pe\n", " attn_out, _ = self.mha(x, x, x, need_weights=False)\n", " return self.fc(attn_out.mean(dim=1)).squeeze(-1)\n", "\n", "def train_eval(use_pos_enc, seed, epochs=300, lr=0.01):\n", " torch.manual_seed(seed)\n", " model = TinyAttnClassifier(use_pos_enc)\n", " opt = torch.optim.Adam(model.parameters(), lr=lr)\n", " for _ in range(epochs):\n", " opt.zero_grad()\n", " loss = F.binary_cross_entropy_with_logits(model(X_train), y_train)\n", " loss.backward()\n", " opt.step()\n", " with torch.no_grad():\n", " return ((model(X_test) > 0).float() == y_test).float().mean().item()\n", "\n", "no_pe_accs = [train_eval(False, seed) for seed in range(5)]\n", "pe_accs = [train_eval(True, seed) for seed in range(5)]\n", "\n", "print(f'without positional encoding: mean test acc = {np.mean(no_pe_accs):.1%} (+/- {np.std(no_pe_accs):.1%})')\n", "print(f'with positional encoding: mean test acc = {np.mean(pe_accs):.1%} (+/- {np.std(pe_accs):.1%})')" ] }, { "cell_type": "markdown", "id": "c1d1aeb2", "metadata": {}, "source": [ "Without positional encoding, the model is stuck at chance — it can only ever report which two tokens showed up, never their relative order, no matter how long it trains. Adding the fixed sinusoid gives every position a distinct signature the model can key off, and the task becomes trivial.\n", "\n", "## The encoder block\n", "\n", "A single Transformer layer is multi-head attention, a small per-position feedforward network, and two residual connections (Lesson 37) with layer normalization (Lesson 37's batch norm, but normalizing across the *feature* dimension for each individual token instead of across the batch):" ] }, { "cell_type": "code", "execution_count": null, "id": "fda8312d", "metadata": {}, "outputs": [], "source": [ "class EncoderBlock(nn.Module):\n", " def __init__(self, D, H, FF):\n", " super().__init__()\n", " self.mha = nn.MultiheadAttention(D, H, batch_first=True)\n", " self.ln1 = nn.LayerNorm(D)\n", " self.ff = nn.Sequential(nn.Linear(D, FF), nn.ReLU(), nn.Linear(FF, D))\n", " self.ln2 = nn.LayerNorm(D)\n", "\n", " def forward(self, x):\n", " attn_out, _ = self.mha(x, x, x, need_weights=False)\n", " x = self.ln1(x + attn_out) # residual + norm around attention\n", " ff_out = self.ff(x)\n", " x = self.ln2(x + ff_out) # residual + norm around the feedforward network\n", " return x\n", "\n", "FF = 32\n", "block = EncoderBlock(D, H, FF)\n", "layer_torch = nn.TransformerEncoderLayer(d_model=D, nhead=H, dim_feedforward=FF,\n", " batch_first=True, dropout=0.0)\n", "layer_torch.eval()\n", "\n", "# copy torch's weights into our block so the two are directly comparable\n", "block.mha.load_state_dict(layer_torch.self_attn.state_dict())\n", "block.ln1.load_state_dict(layer_torch.norm1.state_dict())\n", "block.ln2.load_state_dict(layer_torch.norm2.state_dict())\n", "block.ff[0].load_state_dict({'weight': layer_torch.linear1.weight, 'bias': layer_torch.linear1.bias})\n", "block.ff[2].load_state_dict({'weight': layer_torch.linear2.weight, 'bias': layer_torch.linear2.bias})\n", "\n", "x2 = torch.randn(2, 6, D)\n", "out_block = block(x2)\n", "out_layer = layer_torch(x2)\n", "print(f'encoder block max abs diff vs nn.TransformerEncoderLayer: {(out_block - out_layer).abs().max().item():.2e}')\n", "print(f'output shape unchanged from input: {tuple(out_block.shape)} == {tuple(x2.shape)}')" ] }, { "cell_type": "markdown", "id": "d80e0a78", "metadata": {}, "source": [ "The block's output has exactly the same shape as its input — same trick as Lesson 37's residual block, and for the same reason: it means blocks can be stacked arbitrarily deep, each one refining the same sequence of vectors, without any reshaping between them. A full Transformer *encoder* is just `N` copies of this block stacked in sequence (`nn.TransformerEncoder` in PyTorch); a *decoder* adds a second attention step per block that attends to the encoder's output, plus a causal mask (Lesson 45) on its own self-attention. Vision Transformers (Lesson 47) reuse the encoder side almost unchanged — the only real difference is what gets fed in as the initial sequence of vectors.\n", "\n", "### Exercises\n", "\n", "1. Change `H` (heads) from 4 to 1, keeping `D=16` fixed, and rerun the multi-head attention validation. With a single head, is there still a meaningful difference from Lesson 45's single-head `attention` function?\n", "2. In the order-detection task, change `T` (sequence length) from 8 to 32. Does the with-positional-encoding model's accuracy hold up, or does the longer sequence make the task harder in a way positional encoding alone doesn't fix?\n", "3. `EncoderBlock` above uses \"post-norm\" (`LayerNorm` applied *after* the residual add, matching the original 2017 paper). Many modern Transformers use \"pre-norm\" instead: `x = x + self.mha(self.ln1(x))`. Implement pre-norm and compare the two at greater depth (stack 10 blocks) — does one train more stably, echoing Lesson 37's vanishing-gradient story?" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.10.0" } }, "nbformat": 4, "nbformat_minor": 5 }