Attention
Be able to explain attention as a "weighted sum steered by similarity", work out the attention weights for a small example with Q, K, V, and implement scaled dot-product attention.
Prerequisites
Intuition
"The bank was by the river." To understand bank the model has to look at river. Attention lets every word look at all the other words and weight them: a lot of weight on what is relevant, little on the rest. The result is a new vector for the word — a mixture of all the words' information, weighted by relevance.
Three roles per word: Query (what I am looking for), Key (what I offer), Value (the information I pass on). The similarity between my query and the others' keys gives the weights; the weights mix their values.
Formal
With Q, K, V as matrices (one row per token, d columns):
Attention(Q, K, V) = softmax(QKᵀ / √d) · V
- QKᵀ: the dot product between every query and every key — an n×n similarity matrix.
- /√d: scaling, so the softmax does not get too peaked when d is large.
- softmax row by row: every row becomes weights that sum to 1.
- · V: a weighted sum of the values.
Q, K and V come from the same input X through three learnt matrices: Q = XW_Q, K = XW_K, V = XW_V. It is self-attention when the same sequence supplies all three.
A small example, d = 1: q = 1, keys = (1, 3), values = (10, 20). Scores = (1, 3); softmax ≈ (0.12, 0.88); the output ≈ 0.12·10 + 0.88·20 = 18.8.
Code
import numpy as np
def softmax(z):
e = np.exp(z - z.max(axis=-1, keepdims=True))
return e / e.sum(axis=-1, keepdims=True)
def attention(Q, K, V):
d = Q.shape[-1]
scores = Q @ K.T / np.sqrt(d)
weights = softmax(scores) # (n, n), every row sums to 1
return weights @ V, weights
rng = np.random.default_rng(0)
X = rng.normal(size=(4, 8)) # 4 tokens, 8 dimensions
WQ, WK, WV = (rng.normal(size=(8, 8)) for _ in range(3))
out, w = attention(X @ WQ, X @ WK, X @ WV)
print(w.sum(axis=1)) # [1. 1. 1. 1.]
Mastery means
- Calculates attention weights with a softmax over QKᵀ/√d for a small example
- Implements attention in NumPy/PyTorch and verifies that the weights sum to 1
Sign in to do the exercises and build your mastery up.
Sources
Leads to
Part of the goals (25)
- Build a transformer from scratch
- Understand how generative AI works
- Run models more cheaply: quantisation
- Interpreting a language model
- Language models in practice
- Build a voice interface
- Frontier Lab — an independent research project
- Build a RAG system you can trust
- Evals in practice
- Reproduce a paper
- Multimodal systems
- Fine-tune and run your own models
- AI in production
- AI safety in practice
- Fine-tune a model with LoRA
- Responsible AI in practice
- Build an agent you can trust
- Build an NLP system end to end
- Generative models in depth
- An AI service in operation
- Build a memory system for an agent
- Build an AI service that survives production
- AI, ethics and society
- Deep reinforcement learning
- Statistics for experiments