Attention from scratch in NumPy
Build one complete attention head out of nothing but arrays, then check every number against PyTorch's own kernel.
- 11 min read
- 3 reading levels
- Updated
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
An attention head is seven small steps in a row. You can build all of them yourself in an afternoon.
Think about making chai from a ready-made sachet against making it from scratch. The sachet works, but you have no idea what is inside. Boil the water, add the leaves, crush the ginger, pour the milk. Now you know which step makes it sweet and which makes it strong.
Attention is the same. Every library gives you a one-line version. Build it once by hand and it stops being magic.
The seven steps
Nothing here is new. This lesson takes the pieces from the previous two and lays them end to end.
1. take the word vectors -> X
2. make a query from each one -> Q
3. make a key from each one -> K
4. make a value from each one -> V
5. score every query against every key
6. turn the scores into shares that add to one
7. mix the values in those shares, then reshape the resultSteps two, three and four each use their own set of learned numbers. Steps five, six and seven have no learned numbers at all — they are pure arithmetic.
That split surprises people. The famous part of attention, the scoring and mixing, has nothing to learn. All the learning sits in the four small matrices around it.
Why building it yourself is worth the hour
You will hit three specific confusions. Hitting them on purpose is cheaper than hitting them mid-project.
Which way round do the shapes go? Rows are tokens. Columns are slots. Getting this backwards produces code that runs and results that are nonsense.
Which direction does the sharing happen in? Each row of the score grid belongs to one word doing the looking. The shares must add to one across that row, not down the column.
Where do the learned numbers actually live? In four matrices, not in the scoring step.
How you know you got it right
There is a clean test. PyTorch ships its own attention, written by people who do this professionally. Feed both versions the same inputs and subtract the answers.
If the difference is around a millionth of a millionth of a millionth, you got it right. That tiny leftover is rounding, not error. Computers store fractions with limited room. Two correct routes to one answer differ in the final digits.
Remember this
- An attention head is four learned matrices wrapped around three fixed arithmetic steps.
- The scoring and mixing steps have nothing to learn.
- Always check a hand-built version against a trusted one before trusting it.
What to learn next
- Multi-head attention — run several of these side by side.
- Einsum in PyTorch — a clearer notation for exactly these contractions.
- Tensor shapes and broadcasting — the skill that prevents most attention bugs.
Developer — Code and libraries.
Setup
pip install numpy torchNumPy does the work. PyTorch appears only at the end, as a referee.
One complete head
import numpy as np
rng = np.random.default_rng(7)
np.set_printoptions(precision=3, suppress=True)
d_model, d_head, n_tokens = 8, 4, 5
words = ["chai", "is", "very", "hot", "today"]
X = rng.normal(size=(n_tokens, d_model)) * 0.5 # stand-in for real embeddings
W_q = rng.normal(size=(d_model, d_head)) * 0.5
W_k = rng.normal(size=(d_model, d_head)) * 0.5
W_v = rng.normal(size=(d_model, d_head)) * 0.5
W_o = rng.normal(size=(d_head, d_model)) * 0.5
def softmax(x, axis=-1):
x = x - x.max(axis=axis, keepdims=True)
e = np.exp(x)
return e / e.sum(axis=axis, keepdims=True)
Q = X @ W_q # step 1: ask
K = X @ W_k # step 2: advertise
V = X @ W_v # step 3: offer
logits = Q @ K.T / np.sqrt(d_head) # step 4: match every query to every key
A = softmax(logits) # step 5: turn matches into a mixing recipe
ctx = A @ V # step 6: mix the values
out = ctx @ W_o # step 7: put it back in model width
for name, m in [("X", X), ("Q", Q), ("K", K), ("V", V),
("logits", logits), ("A", A), ("ctx", ctx), ("out", out)]:
print(f"{name:7s} {str(m.shape):10s}")
print("\nattention matrix A (row = the token doing the looking):")
print(" " + "".join(f"{w:>8s}" for w in words))
for w, row in zip(words, A):
print(f"{w:>8s} " + "".join(f"{v:8.3f}" for v in row))
print("\nevery row sums to", A.sum(axis=1))
# Cross-check against PyTorch's built-in kernel.
import torch
q = torch.tensor(Q).unsqueeze(0).unsqueeze(0) # (batch, head, tokens, d_head)
k = torch.tensor(K).unsqueeze(0).unsqueeze(0)
v = torch.tensor(V).unsqueeze(0).unsqueeze(0)
ref = torch.nn.functional.scaled_dot_product_attention(q, k, v)
print("\ntorch version:", torch.__version__)
print("max difference vs torch's kernel:", float((ref.squeeze().numpy() - ctx).__abs__().max()))X (5, 8)
Q (5, 4)
K (5, 4)
V (5, 4)
logits (5, 5)
A (5, 5)
ctx (5, 4)
out (5, 8)
attention matrix A (row = the token doing the looking):
chai is very hot today
chai 0.221 0.248 0.193 0.155 0.184
is 0.190 0.190 0.197 0.188 0.235
very 0.260 0.258 0.197 0.143 0.142
hot 0.245 0.234 0.204 0.175 0.141
today 0.151 0.202 0.198 0.184 0.266
every row sums to [1. 1. 1. 1. 1.]
torch version: 2.5.1+cu121
max difference vs torch's kernel: 2.220446049250313e-16What the shapes are telling you
Follow the widths down the list. X is 8 wide. Q, K and V are 4 wide. ctx is 4 wide, and out is back to 8.
That narrowing is deliberate. A head works in a smaller space than the model. W_o puts its result back into model width, so it can be added to the residual stream. Without that final projection, the shapes would not line up and the head could not be plugged in.
The logits and A matrices are the only square ones, and they are square in tokens, not in width. Five tokens, so twenty-five scores. Ten tokens would be a hundred. This is where the quadratic cost comes from.
The attention matrix is boring here, and that is correct
Every weight sits near 0.2, which is one divided by five. With random untrained weights, that is the right answer. A fresh head has no reason to prefer any token.
Do not read meaning into a random head. Structure appears only after training on real text.
The cross-check is the most important line
2.220446049250313e-16 is machine epsilon for double precision — the smallest gap NumPy can represent near 1.0. The two implementations agree exactly, and the leftover is the last bit of a 64-bit float.
If your number comes back around 1e-7 instead, you are comparing float32 tensors, which is also fine. If it comes back around 0.1, something is genuinely wrong. Check these in order:
- Is the softmax over the last axis?
- Is the division by the square root of
d_head, notd_model? - Is
Ktransposed on the right operand, giving a tokens-by-tokens result?
The same thing in PyTorch, end to end
import torch, torch.nn as nn, torch.nn.functional as F
class Head(nn.Module):
def __init__(self, d_model, d_head):
super().__init__()
self.q = nn.Linear(d_model, d_head, bias=False)
self.k = nn.Linear(d_model, d_head, bias=False)
self.v = nn.Linear(d_model, d_head, bias=False)
self.o = nn.Linear(d_head, d_model, bias=False)
def forward(self, x): # x: (batch, tokens, d_model)
q, k, v = self.q(x), self.k(x), self.v(x)
ctx = F.scaled_dot_product_attention(q.unsqueeze(1), k.unsqueeze(1), v.unsqueeze(1))
return self.o(ctx.squeeze(1))
print(Head(8, 4)(torch.randn(2, 5, 8)).shape)torch.Size([2, 5, 8])
The unsqueeze(1) adds the head axis. scaled_dot_product_attention wants (batch, heads, tokens, width). With more than one head you would reshape instead — see multi-head attention.
Common mistakes
Using np.dot on stacked arrays and getting a shape you did not expect. For arrays with more than two dimensions, np.dot and @ behave differently. Use @, which is matmul, and treat leading dimensions as batch.
Forgetting keepdims=True in the softmax. Without it the max and the sum lose their axis, and broadcasting misaligns. You get a plausible-looking wrong answer with no error message.
Building the head so it outputs head width instead of model width. The output projection is not optional decoration. Without it there is nothing to add back to the residual stream.
Comparing against PyTorch in a different dtype. NumPy defaults to float64, PyTorch to float32. Cast one side deliberately, or you will chase a 1e-7 difference that was never a bug.
Try it yourself
Set W_k = W_q so keys and queries share one matrix. The score matrix becomes symmetric, meaning A[i][j] mirrors A[j][i] before the softmax. Print logits - logits.T and confirm it is all zeros.
Then reason about why real models keep them separate. "Delhi" should attend to "capital" far more than "capital" attends to "Delhi".
What to learn next
- Multi-head attention — run several of these side by side.
- Einsum in PyTorch — a clearer notation for exactly these contractions.
- Tensor shapes and broadcasting — the skill that prevents most attention bugs.
Researcher — Mathematics and papers.
The complete head as one expression
$$ \operatorname{Head}(X) = \operatorname{softmax}!\left( \frac{X W_Q W_K^\top X^\top}{\sqrt{d_k}} \right) X W_V W_O $$
with $X \in \mathbb{R}^{T \times d_{\text{model}}}$, $W_Q, W_K \in \mathbb{R}^{d_{\text{model}} \times d_k}$, $W_V \in \mathbb{R}^{d_{\text{model}} \times d_v}$ and $W_O \in \mathbb{R}^{d_v \times d_{\text{model}}}$.
Written this way, the two products $W_Q W_K^\top$ and $W_V W_O$ appear as units. Elhage et al. (2021), A Mathematical Framework for Transformer Circuits, name them the QK circuit and the OV circuit. Neither factorisation is identifiable on its own. Replace $(W_Q, W_K)$ with $(W_Q M, W_K M^{-\top})$ for any invertible $M \in \mathbb{R}^{d_k \times d_k}$, and the head's function is unchanged. Any claim about an individual projection matrix must therefore be invariant under this reparameterisation, or it is an artefact.
Numerical notes for a hand-rolled implementation
Softmax must subtract the row maximum. Without it, logits above roughly 88 overflow float32. Subtracting a per-row constant is exact. It changes no output value, only the intermediate magnitudes.
Verify against a reference in the same dtype. In float64, agreement should be at $10^{-16}$. In float32, at $10^{-7}$. A discrepancy of $10^{-3}$ or worse is a bug, not accumulation.
Do not ship the explicit form. Materialising the $T \times T$ matrix costs $O(T^2)$ memory per head per layer. F.scaled_dot_product_attention dispatches to FlashAttention or a memory-efficient kernel. Both keep memory linear in $T$ (Dao et al., 2022, arXiv:2205.14135; Rabe and Staats, 2021, arXiv:2112.05682). The explicit version is a teaching tool and a debugging fallback.
Gradients, if you are implementing backward by hand
With $A = \operatorname{softmax}(S)$, $S = QK^\top / \sqrt{d_k}$ and $O = AV$, and writing $\bar{Z}$ for $\partial \mathcal{L} / \partial Z$:
$$ \bar{V} = A^\top \bar{O}, \qquad \bar{A} = \bar{O} V^\top $$
$$ \bar{S}{ij} = A{ij}\left( \bar{A}{ij} - \sum{k} A_{ik} \bar{A}_{ik} \right) $$
$$ \bar{Q} = \frac{\bar{S} K}{\sqrt{d_k}}, \qquad \bar{K} = \frac{\bar{S}^\top Q}{\sqrt{d_k}} $$
The row-wise subtraction in $\bar{S}$ is the softmax Jacobian applied per row. It is also the term that vanishes when $A$ saturates. If $A_{ij} \to 1$ for one $j$ and $0$ elsewhere, every entry of $\bar{S}$ tends to zero. Validate any hand-written backward with torch.autograd.gradcheck in float64 — see verifying gradients with gradcheck.
Complexity
| Step | FLOPs | Peak memory |
|---|---|---|
| $Q, K, V$ projections | $6 T d_{\text{model}} d_k$ | $3 T d_k$ |
| $QK^\top$ | $2 T^2 d_k$ | $T^2$ |
| softmax | $O(T^2)$ | $T^2$ |
| $AV$ | $2 T^2 d_v$ | $T d_v$ |
| output projection | $2 T d_v d_{\text{model}}$ | $T d_{\text{model}}$ |
The projections are linear in $T$; the middle three are quadratic. The crossover sits near $T \approx 2 d_{\text{model}}$.
Papers
- Vaswani et al., Attention Is All You Need, 2017 — arxiv.org/abs/1706.03762
- Elhage et al., A Mathematical Framework for Transformer Circuits, 2021 — transformer-circuits.pub/2021/framework
- Rabe and Staats, Self-attention Does Not Need O(n²) Memory, 2021 — arxiv.org/abs/2112.05682
- Dao et al., FlashAttention, 2022 — arxiv.org/abs/2205.14135
What to learn next
- Multi-head attention — run several of these side by side.
- Einsum in PyTorch — a clearer notation for exactly these contractions.
- Tensor shapes and broadcasting — the skill that prevents most attention bugs.