Trust is earned, not given

A different perspective

2023-08-30 · Projects

AI fundamentals, part 2: Attention — the idea that made transformers win

Part 2from the AI fundamentals series · 8 parts in all

Part 2 of the series. With text turned into token IDs, the model must decide which previous tokens matter for the next one. In "The key to the cabinet is rusted", the word after "rusted" depends on "key", ten words back, not on "The". The mechanism that learns these links is attention — the single idea that carried transformers to dominance. We'll implement it from scratch, in about 40 lines.

The retrieval metaphor: queries, keys, values

Picture a soft database lookup. Every token position emits three vectors, all learned:

Each position's query is scored against every position's key (a dot product — high when the vectors align), the scores are softmaxed into percentages that sum to 1, and the output is the corresponding weighted sum of values. "Looking up" becomes a smooth blend instead of a hard fetch, and it's differentiable — so gradient descent can learn what to look for.

Scaling: why the divide by sqrt(d)

Dot products grow with vector length; a 512-dimension dot product can easily reach ±100, and softmax of big numbers collapses to one-hot (all attention on one token, gradients die). Dividing by sqrt(d_k) keeps the logits in a sane range. It's one line and it matters enormously.

The full self-attention, in NumPy

import numpy as np

def softmax(x, axis=-1):
    # Numerically stable softmax: subtract the max before exponentiating.
    x = x - x.max(axis=axis, keepdims=True)
    e = np.exp(x)
    return e / e.sum(axis=axis, keepdims=True)

def self_attention(X, Wq, Wk, Wv, Wo, causal=True):
    # X:  (T, d_model)  input embeddings for T tokens
    # Wq/Wk/Wv: (d_model, d_k) learned projections
    # Wo: (d_k, d_model) output projection
    # causal: hide the future (language models must not peek ahead)
    T, d_k = X.shape[0], Wq.shape[1]

    Q = X @ Wq                                    # (T, d_k) what each token asks
    K = X @ Wk                                    # (T, d_k) what each token offers
    V = X @ Wv                                    # (T, d_k) what each token gives

    scores = Q @ K.T / np.sqrt(d_k)               # (T, T) raw affinity, scaled

    if causal:                                    # the language-model mask:
        mask = np.triu(np.ones((T, T), bool), 1)  # position i may look at
        scores[mask] = -1e9                       # positions <= i only; banned
                                                  # entries get -inf so softmax
                                                  # gives them exactly 0
    A = softmax(scores)                           # rows now sum to 1.0

    return (A @ V) @ Wo, A                        # blended context, and the
                                                  # attention map for inspection

Run it on toy embeddings and inspect A — row i shows exactly where token i "looked". The mask triangle is what makes the model a language model: when predicting token 50, tokens 51+ simply don't exist yet.

Multi-head: eight opinions beat one

One attention map can only express one kind of relationship. Real transformers run h parallel heads in smaller subspaces (8 heads × 64 dims instead of 1 × 512), then concatenate. Different heads specialize — one tracks syntax, one coreference, one positional patterns — observed by literal inspection of trained maps:

def multi_head_attention(X, params, n_heads=8):
    # params holds Wq/Wk/Wv/Wo for all heads. Splits d_model into n_heads chunks
    # so the total compute stays the same as single-head attention.
    T, d_model = X.shape
    d_head = d_model // n_heads
    heads = []
    for h in range(n_heads):
        sl = slice(h * d_head, (h + 1) * d_head)      # this head's subspace
        out, _ = self_attention(X, params["Wq"][:, sl], params["Wk"][:, sl],
                                params["Wv"][:, sl], params["Wo"][sl, :], causal=True)
        heads.append(out)
    return np.concatenate(heads, axis=1)              # (T, d_model) again

What attention still needs: positions

Attention alone is a bag of tokens — "dog bites man" and "man bites dog" are identical to it. Positional information gets injected into the embeddings before attention ever runs (the classic sine-cosine scheme, learned embeddings, and today's rotary embeddings are all answers to the same question). That's the next part's subject: assembling embeddings + attention + feed-forward layers + residual connections into a complete GPT block — trainable, in pure NumPy.