Words depend on other words
Read the sentence “the cat ate its fish”. To understand its, you have to look back at cat. To understand ate, you need to know who ate (the cat) and what was eaten (the fish). The meaning of a word depends on the other words around it.
Older language models (RNNs) read text one word at a time and had to carry everything in a single memory vector, so information from far back faded away. In 2017 the paper “Attention Is All You Need” introduced the transformer, where every word can look directly at every other word. That mechanism, self-attention, is the heart of ChatGPT, Gemini, Claude, BERT and most modern AI.
Query, key and value
Think of a library. You walk in with a query (“books about cats”). Every book has a key on its spine (its title and subject). You compare your query with every key, pick the books that match best, and read their contents, the values.
Self-attention gives every word three vectors:
- Query (q): what this word is looking for. “its” looks for its owner.
- Key (k): what this word offers to others. “cat” offers “I am a noun, an animal”.
- Value (v): the information this word passes on if someone listens to it.
All three come from the word’s embedding x multiplied by three learned weight matrices: q = x·W_Q, k = x·W_K, v = x·W_V. Training adjusts these matrices so that useful words find each other. In the 3D model the vectors are 2-dimensional and hand-picked so you can follow every number. Real models use hundreds of dimensions.
Step by step for “its”
- Score every word: the dot product q·k tells how well each key answers the query.
- Scale by √d_k (here √2) so the numbers stay moderate.
- Softmax turns the scores into positive weights that add up to 1.
- Mix the values: z = Σ weight × value.
| Word | Key k | q·k with q = [2, 0] | ÷ √2 | Weight |
|---|---|---|---|---|
| the | [0, 0] | 0 | 0.00 | 0.013 |
| cat | [3, 0] | 6 | 4.24 | 0.895 |
| ate | [0, 3] | 0 | 0.00 | 0.013 |
| its | [0.5, 0] | 1 | 0.71 | 0.026 |
| fish | [1, 0.5] | 2 | 1.41 | 0.053 |
“its” pays about 90 % of its attention to “cat”. Its new vector is z = 0.895 × v_cat + 0.053 × v_fish + … ≈ [0.90, 0.06], which is mostly the cat’s information. After this layer, the word “its” knows what it refers to.
The formula
Doing this for every word at once is just matrix maths:
Attention(Q, K, V) = softmax( Q · Kᵀ / √d_k ) · V
Q, K and V are matrices with one row per word. Q · Kᵀ is the n × n table of scores, the same table that fills up on the right of the 3D model. Because it is one big matrix multiplication, GPUs compute it for all words in parallel. That is why transformers train so much faster than RNNs.
The causal mask (GPT) vs full attention (BERT)
- GPT-style models generate text one word at a time, so while training, a word must not peek at the words after it. Their scores are set to −∞ before softmax, which gives them exactly 0 weight. Only the lower triangle of the table remains.
- BERT-style models read a whole sentence to understand it (search, classification), so every word may look at every other word.
With the mask on, “cat” can only see “the” and itself. Neither key matches its query, so its attention is split 50 / 50.
More than one head, more than one layer
- Multi-head attention runs several attentions side by side, each with its own W_Q, W_K, W_V. One head might track “who did what”, another “which noun does this pronoun mean”.
- Positional encoding is added to each embedding, because attention by itself does not know word order.
- A transformer block is attention + a small feed-forward network, each wrapped with a residual connection and layer normalisation. Large models stack dozens of these blocks.
Code
import numpy as np
words = ["the", "cat", "ate", "its", "fish"]
Q = np.array([[1.5, 0], [0, 1.2], [1.2, 0], [2, 0], [0, 1.5]])
K = np.array([[0, 0], [3, 0], [0, 3], [0.5, 0], [1, 0.5]])
V = np.array([[0, 0], [1, 0], [0, 0.5], [0.3, 0], [0, 1]])
def attention(Q, K, V, causal=False):
d_k = Q.shape[1]
scores = Q @ K.T / np.sqrt(d_k) # n x n table of scores
if causal:
n = len(scores)
scores[np.triu_indices(n, k=1)] = -np.inf # hide later words
weights = np.exp(scores - scores.max(axis=1, keepdims=True))
weights /= weights.sum(axis=1, keepdims=True) # softmax, row by row
return weights @ V, weights
Z, W = attention(Q, K, V)
print(np.round(W[3], 3)) # [0.013 0.895 0.013 0.026 0.053] "its" looks at "cat"
print(np.round(Z[3], 3)) # [0.903 0.059]
In PyTorch the whole thing is one call: torch.nn.functional.scaled_dot_product_attention(q, k, v, is_causal=True).
Common mistakes
- Forgetting to scale by √d_k. It works on tiny examples, but training large models becomes unstable.
- Applying softmax over the wrong axis. Each row (one query against all keys) must sum to 1.
- Applying the causal mask after softmax. Masked words must get −∞ before softmax, or the weights no longer sum to 1.
- Treating attention weights as a full explanation of the model’s decision. They show where information flows in one layer, not the whole reasoning.
Complexity at a glance
| Case / operation | Time | Why |
|---|---|---|
| Attention over n tokens (d dimensions) | O(n² · d) | Every token is compared with every other token. |
| Memory for the attention table | O(n²) | Why very long contexts are expensive. |
| Sequential steps per layer | O(1) | All tokens are processed at once (an RNN needs n steps). |
| Extra space | O(n² + n · d) |
Quick check
Test yourself — pick an answer to see if you got it.
1. What does softmax guarantee about each row of attention weights?
Softmax computes e^s / Σ e^s, so every weight is positive and each row sums to 1 — a recipe for how much to listen to each word.
2. Why are the dot-product scores divided by √d_k?
Dot products of long vectors have large values. Scaling by √d_k keeps them in a range where softmax stays smooth and gradients keep flowing.
3. In a GPT-style (causal) model reading “the cat ate its fish”, which words can “cat” attend to?
The causal mask hides every later word, because a model that writes text one word at a time cannot see the future.
4. Why does attention get expensive as the text gets longer?
Doubling the context length makes the attention table four times bigger, in both time and memory.