Attention: Which Position Is Comparing With Which?
Build multi-head attention from tensor geometry, name every query, key, head, and position axis, verify masks, softmax, scaling, and head round-trips, then use four levels of evidence—from shape to behavioral intervention—to locate the first semantic failure.
Here is a tensor of sequence representations and the line that splits it into heads.
B, T, E, Nh = 1, 4, 8, 2
Dh = E // Nh # 4
x = torch.arange(B * T * E).reshape(B, T, E)
heads = x.reshape(B, Nh, T, Dh)
x: (1, 4, 8)
heads: (1, 2, 4, 4)
That is exactly the shape multi-head attention wants: batch, heads, positions, head dimension. Nothing raised. Every shape assertion passes.
Now the same split written the other way:
good = x.reshape(B, T, Nh, Dh).transpose(1, 2)
heads shape (1, 2, 4, 4)
good shape (1, 2, 4, 4)
same shape : True
same dtype : True
same element multiset: True
same values : False
Two tensors, identical shape, identical dtype, containing the same 32 numbers, and they are not the same tensor. Run both through the same attention computation and the outputs differ:
A reshape can change a shape. It cannot perform a transpose you forgot to express.
That is the chapter, stated at the smallest possible scale. Attention is a sequence of explicit relationships between represented vectors, and almost every stage of it can be wrong while producing exactly the shape you predicted.
Where we are
Chapter 8 removed the last of the ambiguity from [B, C, H, W] by deriving every channel and spatial transformation before running the layer, and comparing the derivation with the observation. Chapter 9 then removed the spatial grid entirely and left us with:
[B, T, D]
B how examples are grouped
T positions within one example
D coordinates describing the represented object at one position
Chapter 9 also made one operation concrete. A dot product between two D-vectors is a score along a direction, not automatically a similarity, a cosine or a distance. And it ended by noticing that a learned direction w was never special: it is a vector in the same space as x, so nothing stops one vector in a batch from being scored against another vector in the same batch.
[B, T, D] contains T such vectors per example. Scoring every one against every other produces a T × T grid of numbers, and the question this chapter answers is what happens next:
When attention turns
[B, T, E]into pairwise comparisons, what does every axis mean at every stage, and where did the relationships first stop meaning what we intended?
Notice what that question is not. It is not “what is attention?”. Plenty of people can recite Q, K, V, softmax and multi-head and still lose a day to a model whose head axis holds the wrong values. The useful skill is opening unfamiliar attention code, naming every axis, and saying what should be true after each operation.
The environment
Every shape, number, exception and comparison in this chapter came from executing the code shown, with fixed seeds, on PyTorch 2.13.0 and Python 3.12.3, Linux CPU.
Attention APIs are unusually version-sensitive, and mask conventions in particular differ between functions in the same release. Everything claimed here about F.scaled_dot_product_attention and nn.MultiheadAttention was checked against the installed version’s own documentation and then executed.
Names first
Chapter 8 used H for image height, and the transformer literature uses H for the number of heads. That collision does real damage when you are working out which axis a transpose moved, so this chapter uses explicit names throughout:
B batch, or any leading grouping axis
Tq query positions
Tk key / value positions
E model dimension, the size of one position's representation
Nh number of heads
Dh head dimension for queries and keys
Dv head dimension for values
For the ordinary multi-head self-attention implementation built in this chapter:
E = Nh * Dh
Tq = Tk = T
Dv = Dh
Those are contracts of this implementation, not laws of attention: grouped-query and multi-query attention use fewer key/value heads than query heads, and some architectures use a value dimension different from the key dimension. The derivations below keep Tq, Tk, Dh and Dv separate for that reason, and collapse them only where the specific implementation requires it.
What the bad reshape actually did
“The values get scrambled” is not an explanation. Here is the mechanism, using the tiny deterministic tensor from the opening.
The input is four positions of eight features each:
tensor([[ 0, 1, 2, 3, 4, 5, 6, 7], position 0
[ 8, 9, 10, 11, 12, 13, 14, 15], position 1
[16, 17, 18, 19, 20, 21, 22, 23], position 2
[24, 25, 26, 27, 28, 29, 30, 31]]) position 3
The intent of a head split is to partition the feature axis. Each position’s eight features become two groups of four, one per head:
position 0: head 0 <- [ 0 1 2 3] head 1 <- [ 4 5 6 7]
position 1: head 0 <- [ 8 9 10 11] head 1 <- [12 13 14 15]
position 2: head 0 <- [16 17 18 19] head 1 <- [20 21 22 23]
position 3: head 0 <- [24 25 26 27] head 1 <- [28 29 30 31]
reshape(B, T, Nh, Dh) does exactly that, because it splits the last axis and leaves everything before it alone. transpose(1, 2) then moves the head axis, which now exists, out in front of the position axis:
good[0, 0] (head 0) good[0, 1] (head 1)
tensor([[ 0, 1, 2, 3], tensor([[ 4, 5, 6, 7],
[ 8, 9, 10, 11], [12, 13, 14, 15],
[16, 17, 18, 19], [20, 21, 22, 23],
[24, 25, 26, 27]]) [28, 29, 30, 31]])
Four rows per head, one per position, each holding that head’s slice of that position’s features. That is what [B, Nh, T, Dh] is supposed to mean.
Now reshape(B, Nh, T, Dh). A reshape reads the elements in memory order and refills them into the new shape, so it cuts the flat sequence of 32 numbers into two blocks of 16. Memory here is position-major, so the first 16 numbers are positions 0 and 1 in their entirety:
heads[0, 0] (head 0) heads[0, 1] (head 1)
tensor([[ 0, 1, 2, 3], tensor([[16, 17, 18, 19],
[ 4, 5, 6, 7], [20, 21, 22, 23],
[ 8, 9, 10, 11], [24, 25, 26, 27],
[12, 13, 14, 15]]) [28, 29, 30, 31]])
“Head 0” now contains positions 0 and 1, and “head 1” contains positions 2 and 3. Worse, within each head, one position’s two feature halves have become two separate rows on the Tq axis. The tensor claims four query positions per head; it holds two positions each split in half.
The operation did not split the feature axis at all. It split the sequence. Every subsequent stage runs happily on that: the score matrix compares half-positions with half-positions, softmax normalizes over the wrong candidates, and the output has the right shape.
The general form is worth carrying:
reshapereinterprets the existing element order.transposechanges the element order. Askingreshapefor an axis arrangement that requires a permutation gives you the shape and not the arrangement.
Chapter 2 owns the storage mechanics behind that: stride, contiguity, view versus reshape. This chapter is about what the axes mean, and it will lean on Chapter 2 rather than repeat it. One practical pointer, since it comes up twice below: after a transpose the tensor may be non-contiguous, so use reshape, or make the copy explicit with .contiguous().view(...).
The technique: four levels of correctness
Every chapter has added an investigation method. Chapter 2 asked for the first wrong tensor rather than the first illegal one. Chapter 3 asked where the gradient path stops existing. Chapter 5 asked which structure disagreed about ownership. Chapter 7 asked where a sample stopped satisfying its contract. Chapter 8 asked where derived and observed geometry first diverged. Chapter 9 asked what one vector represents and which axis holds its coordinates.
Chapter 10 adds:
Name every attention axis before the operation runs. Then verify the first invariant that should become true after it. Stop at the first stage where that invariant fails.
Shape alone will not carry you here, and the opening failure is the proof: it produced exactly the derived shape. So before building anything, it is worth being precise about what “correct” means for an attention tensor, because it means four different things at four different depths.
LEVEL 1 SHAPE Does the tensor have the dimensions I derived?
caught by: an assertion on .shape
LEVEL 2 AXIS SEMANTICS Do those dimensions hold the objects I intended?
caught by: deterministic tensors, value comparison,
round-trip properties
LEVEL 3 NUMERICAL Do query rows sum to one over the key axis?
INVARIANTS Are forbidden weights zero? Is everything finite?
caught by: assertions on values, not shapes
LEVEL 4 BEHAVIOR Does changing a future input move a past output?
Does this match a trusted reference implementation?
caught by: interventions and differential tests
Each level is cheap to mistake for the one above it, and the levels are not interchangeable. Chapter 8’s shape ledger deliberately concentrated on Level 1 because its question was convolutional geometry. Attention makes the limit of that evidence impossible to ignore: the opening failure sails through Level 1 and dies at Level 2, while later failures can pass Levels 1 through 3 and still violate the intended computation.
So the loop is Chapter 8’s, widened:
DERIVE THE AXES
↓
EXECUTE ONE STAGE
↓
VERIFY THE INVARIANT THAT STAGE INTRODUCES
↓
STOP AT THE FIRST FAILURE
The rest of this chapter builds attention one stage at a time, names the invariant each stage introduces, and labels each deliberate failure with the earliest level that exposes it. The point is not that every failure has a unique level; it is that no single level is sufficient.
Q, K and V are three projections
Start with the input contract and three learned linear maps.
B, T, E = 2, 5, 12
x = torch.randn(B, T, E)
q_proj = nn.Linear(E, E, bias=False)
k_proj = nn.Linear(E, E, bias=False)
v_proj = nn.Linear(E, E, bias=False)
q, k, v = q_proj(x), k_proj(x), v_proj(x)
x (2, 5, 12)
q (2, 5, 12) k (2, 5, 12) v (2, 5, 12)
Chapter 9 already derived that: nn.Linear maps [..., E_in] -> [..., E_out] and preserves every leading axis, so all 10 position vectors are transformed independently and the shape is unchanged. Nothing attention-specific has happened yet. Q, K and V are three learned representations of the same input, and so far they are interchangeable as far as the tensor system is concerned.
The usual metaphors are worth using once and then dropping. Calling the query “what this position is looking for”, the key “what this position offers” and the value “what this position contributes” is a serviceable mnemonic, not a definition. The mechanical content is:
Q and K produce pairwise compatibility scores
V holds the vectors that get mixed according to those scores
That is enough to derive everything below, and unlike the metaphor it is checkable.
One query against all keys
Before heads, before masks, the core object. Take one head’s worth of queries and keys, with Tq and Tk deliberately different:
Q = torch.randn(2, 3, 4) # [B, Tq, Dh]
K = torch.randn(2, 5, 4) # [B, Tk, Dh]
Derive the product before running it. Matrix multiplication contracts the last axis of the left operand against the second-to-last of the right, so the two axes that must meet are the ones holding the Dh coordinates. K has Dh last, so it has to move:
K [B, Tk, Dh] observed (2, 5, 4)
K.transpose(-2, -1) [B, Dh, Tk] observed (2, 4, 5)
Q @ Kᵀ [B, Tq, Dh] @ [B, Dh, Tk] -> [B, Tq, Tk] observed (2, 3, 5)
This is the new object Chapter 10 introduces, and it deserves to be read out loud one entry at a time:
scores[b, i, j]
b which example
i which query position is asking
j which key position is being scored
One scalar is the compatibility between query position i and key position j in example b. That is not an interpretation; it is arithmetic, and it can be checked: Chapter 10 compares Tq vectors against Tk vectors and gets Tq × Tk scalars. That is the entire structural difference, and it is why the score matrix is rectangular in general and square only when the query and key sequences happen to be the same one. In ordinary self-attention, where Tq = Tk = T, the explicit score grid contains T² entries per example per head, so doubling sequence length quadruples that grid. What to do about that belongs to the performance chapter; the shape is what matters here.
Which two axes must move
transpose(-2, -1) is not an attention incantation. It is the claim that the last two axes hold key positions and head coordinates, and those are the two that must swap. Once a tensor is [B, Nh, Tk, Dh], k.transpose(1, 2) swaps heads with key positions instead:
k (2, 3, 5, 4) axes [B, Nh, Tk, Dh]
k.transpose(-2, -1) (2, 3, 4, 5) axes [B, Nh, Dh, Tk] <- what we want
k.transpose(1, 2) (2, 5, 3, 4) axes [B, Tk, Nh, Dh] <- heads and keys swapped
q @ k.transpose(1, 2)
RuntimeError: The size of tensor a (3) must match the size of tensor b (5)
at non-singleton dimension 1
Here it raises, which is the lucky case. When Nh, Tk and Dh happen to be equal, nothing raises:
Nh = Tk = Dh = 4
good (2, 4, 4, 4) bad (2, 4, 4, 4) shapes equal: True
values equal: False max abs diff: 14.6166
Level 1 cannot see it. A value comparison against a correctly-derived reference can.
Why sqrt(Dh) scaling exists
Chapter 9 ended by noting that a dot product’s magnitude depends on both vectors’ norms, and that with component scale held roughly fixed, it can grow systematically as the number of coordinates grows. Attention runs into that immediately, because Dh is a design choice that changes between models.
The argument is short. If Q and K components are independent with mean 0 and variance 1, each term q_d k_d in the dot product has mean 0 and variance 1, so summing Dh of them gives a variance of about Dh and a standard deviation of about sqrt(Dh). Measured, with 64 × 32 query and key vectors at each dimension:
Dh raw std scaled std sqrt(Dh) raw maxp scaled maxp raw H scaled H
8 2.8247 0.9987 2.8284 0.4889 0.1604 1.6781 3.0271
32 5.6665 1.0017 5.6569 0.7355 0.1628 0.7572 3.0186
128 11.3023 0.9990 11.3137 0.8710 0.1656 0.3323 3.0103
512 22.6109 0.9993 22.6274 0.9346 0.1645 0.1606 3.0138
The unscaled standard deviation tracks sqrt(Dh) closely across a 64× range of head dimensions. Dividing by sqrt(Dh) holds it at approximately 1.
The two right-hand pairs of columns show why anyone cares. maxp is the mean largest softmax probability per query row; H is the mean softmax entropy over 32 keys, where a uniform distribution would give ln(32) = 3.4657. Unscaled, the largest probability climbs from 0.49 to 0.93 and the entropy collapses from 1.68 to 0.16 purely because the head dimension grew. Scaled, both stay flat.
The causal chain, with its assumption attached:
larger Dh
↓ assuming roughly unit-variance, independent components
larger unscaled dot-product variance
↓
larger differences between the logits in one query row
↓
sharper softmax
Two cautions. The sqrt(Dh) relationship is a consequence of the variance assumption in this experiment, not a guarantee about arbitrary learned representations, whose components are neither unit-variance nor independent in general. And low entropy is not a defect; a model that has learned to attend sharply is doing its job. The experiment says something narrower and more useful: without the scale, softmax sharpness varies with a shape parameter rather than with what the model learned.
Hold on to this one, because it is the chapter’s Level 4 failure. A missing scale breaks no invariant that can be checked locally. It is caught only by comparison against a reference.
Softmax introduces the chapter’s most useful invariant
Scores are unbounded real numbers. Attention turns each query’s row of them into mixing coefficients with a softmax over the key axis:
weights = torch.softmax(scores, dim=-1)
For scores shaped [B, Nh, Tq, Tk], the last axis is Tk, so this normalizes across candidate keys, once for every (batch, head, query) triple. Which gives the invariant this chapter leans on hardest:
Each valid query row becomes a distribution over allowed key positions.
And therefore a one-line diagnostic:
weights.sum(dim=-1) # should be 1 for every valid row
Anchor: the wrong softmax axis — a Level 3 failure
softmax takes a dim and will accept the wrong one. Chapter 9 made the same point about F.normalize; here the consequence is larger.
good = torch.softmax(scores, dim=-1)
bad = torch.softmax(scores, dim=-2)
shapes equal: True (1, 2, 4, 4) both finite: True values equal: False
good.sum(dim=-1) max abs error from 1: 1.19e-07
bad.sum(dim=-1) : [0.8446, 1.1353, 0.7004, 1.3197]
bad.sum(dim=-2) max abs error from 1: 5.96e-08
Two legal operations, identical shapes, all values finite, both perfectly well-defined. They normalize different relationships:
dim=-1 for each query, distribute weight across the keys
dim=-2 for each key, distribute weight across the queries
The second is a coherent computation. It is not attention as intended here, and no shape assertion anywhere in the model can tell the difference. The row-sum check separates them in one line, which is exactly the jump from Level 1 to Level 3.
Values turn weights into an output vector
Now the last piece. With weights over keys and value vectors indexed by keys:
weights [B, Nh, Tq, Tk]
V [B, Nh, Tk, Dv]
weights @ V -> [B, Nh, Tq, Dv]
The contracted axis is Tk, which is exactly right: the key positions are what the weights are distributed over, and they disappear in the sum. One output vector out[b, h, i] is a weighted combination of the Tk value vectors, using query i’s weights.
Two value vectors and hand-chosen weights make that physical:
V (1, 2, 2) [[1, 0], [0, 1]]
weights (1, 1, 2) [[0.75, 0.25]]
out (1, 1, 2) [[0.75, 0.25]]
And with real softmax weights over three keys in five dimensions, comparing the matmul against the weighted sum written out by hand:
weights[0,1]: [0.1721, 0.4352, 0.3927]
out[0,1] : [-0.1035, 0.1147, 0.7471, 0.5384, 0.1653]
manual sum : [-0.1035, 0.1147, 0.7471, 0.5384, 0.1653] match: True
every output coordinate within the per-coordinate min/max of V: True
That last line names a real property. After an ordinary softmax the weights are nonnegative and sum to one, so each output is a convex combination of the value vectors for that query and cannot land outside their per-coordinate minima and maxima. Additive score biases and masks applied before softmax do not change that property as long as the row retains at least one finite, allowed key. The property stops holding if the coefficients are altered after normalization so that they become negative or no longer sum to one.
The whole mechanism now fits in three lines:
score the keys against the query
normalize those scores over the keys
use the normalized scores to mix the values
Everything else in attention is bookkeeping around that.
Now the heads
Heads are not a new kind of object. They are one more structure axis, and they exist because splitting E into Nh groups lets the model run Nh independent score matrices over Nh different Dh-dimensional subspaces instead of one large one.
[B, T, E]
↓ split the feature axis: E = Nh * Dh
[B, T, Nh, Dh]
↓ move the head axis outward
[B, Nh, T, Dh]
Written as a pair of functions:
def split_heads(x, num_heads):
B, T, E = x.shape
assert E % num_heads == 0, f"E={E} not divisible by num_heads={num_heads}"
Dh = E // num_heads
return x.reshape(B, T, num_heads, Dh).transpose(1, 2)
def merge_heads(x):
B, Nh, T, Dh = x.shape
return x.transpose(1, 2).reshape(B, T, Nh * Dh)
Every stage after the split carries the extra leading axis and is otherwise unchanged from the single-head derivation:
Q [B, Nh, Tq, Dh]
K [B, Nh, Tk, Dh]
V [B, Nh, Tk, Dv]
scores [B, Nh, Tq, Tk]
weights [B, Nh, Tq, Tk]
context [B, Nh, Tq, Dv]
Anchor: the split/merge round trip — a Level 2 failure
Here is the check that would have caught the opening failure in one line, before any attention ran. Splitting and merging heads is pure rearrangement: it must return the original tensor exactly.
With a deterministic `[2, 5, 12]` tensor, across every head count that divides 12:
```text
Nh=1 split (2, 1, 5, 12) merge (2, 5, 12) round trip exact: True
Nh=2 split (2, 2, 5, 6) merge (2, 5, 12) round trip exact: True
Nh=3 split (2, 3, 5, 4) merge (2, 5, 12) round trip exact: True
Nh=4 split (2, 4, 5, 3) merge (2, 5, 12) round trip exact: True
Nh=6 split (2, 6, 5, 2) merge (2, 5, 12) round trip exact: True
Nh=12 split (2, 12, 5, 1) merge (2, 5, 12) round trip exact: True
Use a deterministic tensor such as arange, not random data, and compare values rather than shapes. A shape comparison passes for both the correct and the broken split; a value comparison does not. This is Level 2 in its purest form, and it is the cheapest high-value test in the chapter.
Before testing attention, prove that splitting and merging heads is a lossless rearrangement.
The property is unusually valuable because it fails for a whole family of implementations at once: a missing transpose on the way in, a missing transpose on the way out, a transpose of the wrong pair of axes, or a head count that does not match the one used to merge.
Merging is the same mistake in reverse
Going back from [B, Nh, T, Dh] to [B, T, E] requires the head axis to move back inside before the last two axes are joined:
[B, Nh, T, Dh]
↓ transpose(1, 2)
[B, T, Nh, Dh]
↓ reshape the final two axes into one
[B, T, Nh*Dh]
Skip the transpose and reshape straight from [B, Nh, T, Dh], and you get the right shape and the wrong tensor:
wrong shape: (2, 5, 12) equals x: False
x[0,0] : [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
wrong[0,0] : [0, 1, 2, 3, 12, 13, 14, 15, 24, 25, 26, 27]
Position 0’s twelve features have been rebuilt from three different positions’ head slices. The round-trip assertion catches this too, which is why it is worth writing once rather than reasoning about each function separately.
Masks are relations
A mask answers a question about pairs, and the question differs depending on which mask you mean. That distinction is the source of most mask bugs, so it comes before any code.
CAUSAL / ATTENTION RELATION KEY PADDING
Can query i use key j? Is key j real data for example b?
natural axes [Tq, Tk] natural axes [B, Tk]
Both eventually influence a [B, Nh, Tq, Tk] score tensor. That does not make them the same object, and combining them means deliberately broadcasting two different relations onto the same grid.
Causality, read off the score matrix
Once scores[b, h, i, j] means “query i against key j”, autoregressive causality is a statement about indices:
query i may use keys j <= i
query i must not use keys j > i
Draw it with the axes labeled, because an unlabeled triangle tells you nothing about which corner is the future:
allowed (rows = queries, cols = keys) blocked = ~allowed
k0 k1 k2 k3 k4 k0 k1 k2 k3 k4
q0 1 0 0 0 0 q0 0 1 1 1 1
q1 1 1 0 0 0 q1 0 0 1 1 1
q2 1 1 1 0 0 q2 0 0 0 1 1
q3 1 1 1 1 0 q3 0 0 0 0 1
q4 1 1 1 1 1 q4 0 0 0 0 0
allowed keys per query: [1, 2, 3, 4, 5]
The upper-right region is the future because columns are keys, so a column index j greater than the row index i means a key that comes after the query. In code, the allow relation is the lower triangle including the diagonal, and torch.triu(ones, diagonal=1) is verifiably its exact complement:
allowed = torch.ones(T, T, dtype=torch.bool).tril(diagonal=0)
blocked = ~allowed
triu(ones, diagonal=1) equals blocked: True
Keep the variable named for what it means. mask is the name that causes the next problem.
Anchor: mask polarity — a Level 3 or Level 4 failure, depending on the API
Two PyTorch attention APIs in the same release take boolean masks with opposite polarity. This is documented, it is not a bug, and it will still cost you an afternoon.
From the installed F.scaled_dot_product_attention documentation: a boolean attn_mask value of True indicates the element should take part in attention. From the installed nn.MultiheadAttention documentation: for a binary attn_mask, True indicates the position is not allowed to attend, and for key_padding_mask, True indicates the key will be ignored.
The same tensor means opposite things depending on which function receives it. Executed, with identical Q, K and V and a T=4 causal relation:
SDPA is_causal=True vs allow-mask : allclose True
SDPA is_causal=True vs block-mask : allclose False
max abs diff with inverted polarity: 3.4084601
inverted-polarity output finite: True shape (2, 2, 4, 4)
Through SDPA the wrong polarity produces a finite tensor of the correct shape that differs materially from the causal reference, and no local row-sum invariant fails: it is still a valid softmax over some set of keys. A reference comparison or a causal intervention exposes the semantic error. That is Level 4.
Through nn.MultiheadAttention the same mistake surfaces one level earlier. Here are the returned per-head weights both ways:
blocked-polarity mask (correct for MHA), weights[0,0]:
tensor([[1.0000, 0.0000, 0.0000, 0.0000],
[0.5173, 0.4827, 0.0000, 0.0000],
[0.2602, 0.3738, 0.3661, 0.0000],
[0.2756, 0.2385, 0.2327, 0.2532]])
allow-polarity mask (wrong for MHA), weights[0,0]:
tensor([[0.0000, 0.3521, 0.3015, 0.3464],
[0.0000, 0.0000, 0.4802, 0.5198],
[0.0000, 0.0000, 0.0000, 1.0000],
[ nan, nan, nan, nan]])
upper-triangle weight max, block polarity: 0.0
upper-triangle weight max, allow polarity: 1.0
The second matrix is the failure in full. Query 0 attends only to the future. Query 3 has every key blocked, has no valid distribution to normalize, and comes back as NaN — a Level 3 catch. Maximum weight on a forbidden future key is 1.0: the model is reading tokens it must never see.
One practical note while those weights are in view. nn.MultiheadAttention averages attention weights over heads by default; pass average_attn_weights=False to get them per head as above. An average can look entirely reasonable while one head is doing something pathological, and it destroys the per-head structure you need to check that forbidden weights are zero head by head.
Make the convention part of the name of whatever converts between them, so a call site reads as a claim about the destination API rather than about the tensor:
def allowed_to_sdpa_mask(allowed): # SDPA: True participates
return allowed
def allowed_to_mha_mask(allowed): # MHA: True is blocked
return ~allowed
A boolean mask has no universal meaning. Its meaning belongs to the API receiving it.
A mask diagnostic that prints only mask.shape is therefore not a diagnostic. The things that can be right or wrong independently are:
shape does it broadcast onto [B, Nh, Tq, Tk]?
dtype boolean relation, or additive float bias?
device same as the scores?
polarity does True mean allowed or blocked, for this exact function?
counts how many pairs are allowed, and does any query have none?
A mask can have the correct shape, dtype and device, the wrong polarity, and the program still runs to completion.
Broadcasting a padding mask, explicitly
A key padding mask arrives with the axes it naturally has, and must be told which axes of the score grid it is constant along:
[B, Tk]
↓ it varies with neither head nor query position
[B, 1, 1, Tk]
↓ broadcasts against
[B, Nh, Tq, Tk]
pad_valid = torch.tensor([[True, True, True, True],
[True, True, False, False]]) # True = real token
expanded = pad_valid[:, None, None, :]
combined = expanded & allowed # both in ALLOW polarity
pad_valid (2, 4) expanded (2, 1, 1, 4)
allowed (4, 4) combined (2, 1, 4, 4)
masked scores (2, 2, 4, 4)
row sums max error from 1 : 5.96e-08
weight on padded keys, example 1 : 0.0
allowed keys per query, example 1: [1, 2, 2, 2]
Two things earn their place there. Combining the relations with & is only correct because both are in the same polarity, which is why the variables are named allowed and pad_valid rather than mask and mask2. And example 1’s allowed-key counts stop growing at 2 because positions 2 and 3 are padding: both relations are being respected, and the counts prove it.
Chapter 2 established the broadcasting rules. The rule for attention is narrower:
When the axes have semantics, insert the broadcast axes yourself. Do not let a
[B, Tk]tensor find its own alignment against a four-axis score grid.
Fully masked rows
If every key is forbidden for some query, there is no distribution to normalize. What happens then is an implementation detail, and it differs between the manual pipeline and PyTorch’s API. For a manual masked_fill(-inf) followed by softmax:
manual masked scores row 1: [-inf, -inf, -inf]
manual softmax row 1 : [nan, nan, nan]
manual weights all finite : False
row sums : [1.0, nan, 1.0]
The NaN then propagates through weights @ V and contaminates the rest of the model. The same relation through F.scaled_dot_product_attention on this build behaves differently:
SDPA output finite : True
SDPA output row 1 : [0.0, 0.0, 0.0, 0.0]
non-finite count in SDPA output: 0
Do not generalize either result. The manual one is a property of softmax applied to an all--inf row; the SDPA one is observed behavior for this API, version and backend. A fully masked query is therefore something your model contract must account for, not something whose backend behavior you should guess.
If every query is required to produce a meaningful distribution, reject fully blocked rows explicitly:
Now assemble it. The point of writing this by hand is not that PyTorch’s version is inadequate; it is that the intermediates have to be inspectable.
def manual_attention(q, k, v, allow_mask=None, scale=None):
"""q [.., Tq, Dh] k [.., Tk, Dh] v [.., Tk, Dv]
allow_mask: boolean, True = this (query, key) pair may participate."""
Dh = q.shape[-1]
scale = 1.0 / math.sqrt(Dh) if scale is None else scale
scores = q @ k.transpose(-2, -1) * scale
masked = scores if allow_mask is None else scores.masked_fill(~allow_mask, float("-inf"))
weights = torch.softmax(masked, dim=-1)
out = weights @ v
return out, {"scores": scores, "masked_scores": masked, "weights": weights}
Two decisions are deliberate. The mask argument is named allow_mask, so a reader of a call site knows the polarity without opening the function. And the intermediates come back in a dict, because when a comparison fails you need to know which stage diverged, not just that the outputs differ.
Note the order: mask the scores, then normalize. Zeroing forbidden weights after softmax leaves the surviving row summing to less than one, breaking the invariant the row-sum check exists to protect.
The Level 4 test: compare against the trusted implementation
With identical Q, K, V, the same scale, an equivalent mask in SDPA’s polarity, and dropout disabled:
mine, stages = manual_attention(q, k, v, allow_mask=allowed)
theirs = F.scaled_dot_product_attention(q, k, v, attn_mask=allowed, dropout_p=0.0)
torch.testing.assert_close(mine, theirs)
unmasked max abs diff : 2.38e-07 assert_close passed
causal max abs diff, allow mask : 1.19e-07 assert_close passed
causal max abs diff, is_causal : 1.19e-07 assert_close passed
forbidden weight max : 0.0
Both routes to a causal relation agree with the manual implementation and with each other. That is a real result: the head layout, the transpose, the scale, the mask polarity, the softmax axis and the value aggregation are all consistent with the reference.
Now remove the scale and watch what survives:
unscaled vs SDPA max abs diff : 0.7033
unscaled rows still sum to 1 : 1.19e-07
Right shape, finite values, rows summing to one. Every Level 1, 2 and 3 check in this chapter passes. Only the differential comparison catches it, which is the whole argument for having Level 4 in the toolkit.
The controls that make the comparison meaningful matter, because a mismatch caused by an uncontrolled variable teaches nothing:
same Q, K, V tensors same dtype and device
same scale dropout disabled on both sides
one equivalent mask, converted to each API's polarity
If the outputs disagree, compare the intermediates in order rather than guessing: scores, then masked scores, then weights, then output. The first stage that disagrees is the diagnosis.
The multi-head self-attention module
With every stage verified, the module is short.
in (2, 8, 32) -> out (2, 8, 32) finite True
in (1, 1, 32) -> out (1, 1, 32) finite True
in (3, 17, 32) -> out (3, 17, 32) finite True
in (5, 2, 32) -> out (5, 2, 32) finite True
MultiHeadSelfAttention(embed_dim=10, num_heads=3)
ValueError: embed_dim=10 not divisible by num_heads=3
Sequence length 1 and batch size 1 are in that list deliberately: attention code that squeezes an axis somewhere tends to survive T=8 and fail at T=1, and finding that at step 80,000 is worse than finding it here. The divisibility check belongs in the constructor rather than in a reshape forty lines later, for the same reason Chapter 5 gave for declaring structure explicitly: a contract violated at construction time produces a message that names the actual problem.
Two lines in that forward are doing quiet work
The first is dropout_p=self.dropout if self.training else 0.0. The functional API takes a dropout_p and applies it; it does not inspect an enclosing module’s .training flag, because it has no enclosing module — it is a function. The installed documentation says so directly and recommends exactly this pattern. Two modules, one following the advice and one not, both called after .eval():
NaiveDropoutAttn eval() two calls identical: False
CorrectDropoutAttn eval() two calls identical: True
The naive version is still dropping attention weights at evaluation time, so its predictions are stochastic, its validation metric is noisy, and it disagrees with itself between runs. model.eval() did what it always does, which is set a flag. Nothing rewrote the argument being passed to a function call, and nothing was going to.
The second is assert y.shape == x.shape, which is worth keeping and worth being honest about. That module preserves [B, T, E]. So does a version with the head axis scrambled, a version that normalizes over queries, a version with inverted mask polarity, and a version with no scale at all.
Attention can preserve its external shape while being wrong at nearly every internal stage.
The assertion reaches Level 1. That is why an attention module which asserts its output shape and nothing else feels safe and is not.
And the external contract itself can be misread
nn.MultiheadAttention defaults to batch_first=False, meaning [T, B, E]. Pass it a [B, T, E] tensor:
batch_first=True input (2, 5, 8) -> out (2, 5, 8) weights (2, 5, 5)
batch_first=False input (2, 5, 8) -> out (2, 5, 8) weights (5, 2, 2)
The output shape is identical. The module read the first axis as the sequence and the second as the batch, giving 2 positions across 5 independent examples instead of 5 positions across 2. The only visible evidence is the attention-weight shape, which went from [B, Tq, Tk] to [T, B, B]. Level 1 passes at the boundary; Level 2 fails inside.
Write the expected external layout down at every module boundary. Never infer it from habit.
Testing causality by intervention
A triangular mask is one representation of causality. The contract it exists to enforce is a statement about behavior:
The output at position
imust not depend on the input at any positionj > i.
That is directly testable, and much stronger than looking at the mask. Take a causal module in eval mode with dropout disabled, run a sequence, change only the positions after i, and rerun:
i = 2
x2 = x.clone()
x2[:, i + 1:] = torch.randn(B, T - i - 1, E)
y1 = attn(x, is_causal=True)
y2 = attn(x2, is_causal=True)
torch.testing.assert_close(y1[:, :i + 1], y2[:, :i + 1])
changed positions: [3, 4, 5] inputs identical up to i: True
max |y1-y2| over positions 0..i : 0.0
max |y1-y2| over positions i+1.. : 0.3916566
assert_close passed for the causal prefix
Exactly zero on the prefix, clearly nonzero afterwards. The second number matters as much as the first: if both were zero, the test would be passing for the trivial reason that the intervention did nothing.
Two controls confirm the test can fail. The same module with causal masking disabled, and then the same module given a mask with the correct triangular structure but the wrong polarity for this API:
causal masking disabled, max |y1-y2| over positions 0..i: 0.3799190
blocked-polarity mask, max |y1-y2| over positions 0..i: 0.6628050
The second row is the mask-polarity anchor arriving at Level 4. The mask is triangular. The picture is right. Information flows backward from the future anyway, and no amount of staring at torch.triu would have found it. One intervention did.
For a manual implementation you can also check the weight matrix directly:
forbidden = stages["weights"][..., ~allowed].abs().max() # observed 0.0
That is genuine evidence, and it is Level 3 rather than Level 4, because it tests the internal coefficient matrix rather than the input-output behavior. When both are available, prefer the intervention.
A mask that looks causal is a claim. A past output that does not move when the future changes is evidence.
Cross-attention breaks the square
Self-attention uses one sequence for queries, keys and values, so Tq == Tk and the score matrix is square. That coincidence gets absorbed into the mental model, and then cross-attention arrives and the assumptions fall over. Derive it first:
Q [B, Nh, Tq, Dh] = (2, 4, 7, 16)
K [B, Nh, Tk, Dh] = (2, 4, 11, 16)
V [B, Nh, Tk, Dv] = (2, 4, 11, 5)
scores [B, Nh, Tq, Tk] -> (2, 4, 7, 11)
weights [B, Nh, Tq, Tk] -> (2, 4, 7, 11)
out [B, Nh, Tq, Dv] -> (2, 4, 7, 5)
Then run it:
derived scores (2, 4, 7, 11) observed (2, 4, 7, 11)
derived out (2, 4, 7, 5) observed (2, 4, 7, 5)
row sums max error: 1.19e-07
Dv differs from Dh here and nothing objects, because Dh is contracted away in Q @ Kᵀ while Dv rides through the value aggregation. The two are constrained by different things: Dh must match between Q and K, and Tk must match between K and V.
The mask consequence is immediate. A square causal mask does not fit:
RuntimeError: The size of tensor a (7) must match the size of tensor b (11)
at non-singleton dimension 3
A key padding mask does, because its natural axes were [B, Tk] all along:
cross padding mask (2, 11) -> (2, 1, 1, 11)
broadcasts against (2, 4, 7, 11)
padded-key weight max, example 1: 0.0
So the definition to carry is not T × T:
A score matrix is query positions × key positions. Self-attention is the special case where those are the same sequence.
The attention ledger
Chapter 7 had a preprocessing stage report, Chapter 8 a shape ledger, Chapter 9 a feature-space inspector. Attention needs one too, and it must record more than shapes, because the most dangerous failures in this chapter preserve the expected shape. Each row therefore carries three things: the shape, what the axes mean, and the invariant that becomes checkable at that point.
def attention_ledger(x, num_heads, projections, allow_mask=None, splitter=split_heads):
q_proj, k_proj, v_proj, out_proj = projections
B, T, E = x.shape
Nh, Dh = num_heads, E // num_heads
first = None
def row(name, shape, axes, *checks):
nonlocal first
print(f" {name:<14} {str(tuple(shape)):<16} {axes:<16} "
+ "; ".join(text for text, _ in checks))
for text, ok in checks:
if not ok and first is None:
first = name
print(f" {'':<14} {'':<16} {'':<16} ^^ FIRST DIVERGENCE: {text}")
break
print(f" ATTENTION LEDGER B={B} T={T} E={E} Nh={Nh} Dh={Dh}")
row("input", x.shape, "[B,T,E]", (f"E == Nh*Dh {E == Nh*Dh}", E == Nh * Dh))
rt = merge_heads(splitter(x, Nh))
row("head roundtrip", rt.shape, "[B,T,E]",
(f"merge(split(x)) == x {torch.equal(rt, x)}", torch.equal(rt, x)))
q, k, v = (splitter(p(x), Nh) for p in (q_proj, k_proj, v_proj))
row("Q", q.shape, "[B,Nh,Tq,Dh]")
row("K", k.shape, "[B,Nh,Tk,Dh]",
(f"Dh matches Q {q.shape[-1] == k.shape[-1]}", q.shape[-1] == k.shape[-1]))
row("V", v.shape, "[B,Nh,Tk,Dv]",
(f"Tk matches K {k.shape[-2] == v.shape[-2]}", k.shape[-2] == v.shape[-2]))
scores = q @ k.transpose(-2, -1) / math.sqrt(Dh)
fin = bool(torch.isfinite(scores).all())
row("scores", scores.shape, "[B,Nh,Tq,Tk]",
(f"finite {fin}", fin), (f"std {scores.std():.3f}", True))
if allow_mask is None:
masked = scores
row("masked_scores", masked.shape, "[B,Nh,Tq,Tk]", ("no mask", True))
else:
masked = scores.masked_fill(~allow_mask, float("-inf"))
dead = bool((~allow_mask).all(dim=-1).any())
row("masked_scores", masked.shape, "[B,Nh,Tq,Tk]",
(f"allowed {int(allow_mask.sum())}/{allow_mask.numel()}", True),
(f"no fully blocked query {not dead}", not dead))
weights = torch.softmax(masked, dim=-1)
err = (weights.sum(-1) - 1).abs().max().item()
blocked = None if allow_mask is None else (~allow_mask).expand_as(weights)
forbidden = 0.0 if blocked is None or not blocked.any() \
else weights[blocked].abs().max().item()
fin = bool(torch.isfinite(weights).all())
forbidden_ok = math.isfinite(forbidden) and forbidden <= 1e-6
row("weights", weights.shape, "[B,Nh,Tq,Tk]",
(f"key-row sum error {err:.1e}", math.isfinite(err) and err < 1e-5),
(f"forbidden weight max {forbidden:.1e}", forbidden_ok),
(f"finite {fin}", fin))
context = weights @ v
row("context", context.shape, "[B,Nh,Tq,Dv]",
(f"finite {bool(torch.isfinite(context).all())}", bool(torch.isfinite(context).all())))
row("merged", merge_heads(context).shape, "[B,Tq,Nh*Dv]")
out = out_proj(merge_heads(context))
row("output", out.shape, "[B,Tq,E]",
(f"external layout preserved {out.shape == x.shape}", out.shape == x.shape))
print(" first divergence:", first or "none")
return out
On a healthy causal setup:
ATTENTION LEDGER B=2 T=5 E=12 Nh=3 Dh=4
input (2, 5, 12) [B,T,E] E == Nh*Dh True
head roundtrip (2, 5, 12) [B,T,E] merge(split(x)) == x True
Q (2, 3, 5, 4) [B,Nh,Tq,Dh]
K (2, 3, 5, 4) [B,Nh,Tk,Dh] Dh matches Q True
V (2, 3, 5, 4) [B,Nh,Tk,Dv] Tk matches K True
scores (2, 3, 5, 5) [B,Nh,Tq,Tk] finite True; std 0.269
masked_scores (2, 3, 5, 5) [B,Nh,Tq,Tk] allowed 15/25; no fully blocked query True
weights (2, 3, 5, 5) [B,Nh,Tq,Tk] key-row sum error 1.2e-07; forbidden weight max 0.0e+00; finite True
context (2, 3, 5, 4) [B,Nh,Tq,Dv] finite True
merged (2, 5, 12) [B,Tq,Nh*Dv]
output (2, 5, 12) [B,Tq,E] external layout preserved True
first divergence: none
Now the opening bug, with the head split done by reshape alone:
ATTENTION LEDGER B=2 T=5 E=12 Nh=3 Dh=4
input (2, 5, 12) [B,T,E] E == Nh*Dh True
head roundtrip (2, 5, 12) [B,T,E] merge(split(x)) == x False
^^ FIRST DIVERGENCE: merge(split(x)) == x False
Q (2, 3, 5, 4) [B,Nh,Tq,Dh]
K (2, 3, 5, 4) [B,Nh,Tk,Dh] Dh matches Q True
V (2, 3, 5, 4) [B,Nh,Tk,Dv] Tk matches K True
scores (2, 3, 5, 5) [B,Nh,Tq,Tk] finite True; std 0.267
masked_scores (2, 3, 5, 5) [B,Nh,Tq,Tk] allowed 15/25; no fully blocked query True
weights (2, 3, 5, 5) [B,Nh,Tq,Tk] key-row sum error 1.2e-07; forbidden weight max 0.0e+00; finite True
context (2, 3, 5, 4) [B,Nh,Tq,Dv] finite True
merged (2, 5, 12) [B,Tq,Nh*Dv]
output (2, 5, 12) [B,Tq,E] external layout preserved True
first divergence: head roundtrip
Look at what the rest of the ledger says. Every shape is right. Scores are finite with a perfectly ordinary standard deviation. Rows sum to one. Forbidden weights are zero. The output preserves the external layout. Ten of eleven rows report health, and the model is attending over half-positions. That is Level 2 catching something the other ten rows structurally cannot see.
A mask with one query blocked from every key:
masked_scores (2, 3, 5, 5) [B,Nh,Tq,Tk] allowed 11/25; no fully blocked query False
^^ FIRST DIVERGENCE: no fully blocked query False
weights (2, 3, 5, 5) [B,Nh,Tq,Tk] key-row sum error nan; forbidden weight max nan; finite False
context (2, 3, 5, 4) [B,Nh,Tq,Dv] finite False
first divergence: masked_scores
The weights and context rows are non-finite too, and neither is the diagnosis. The mask is. Same discipline as Chapter 8’s first-divergence rule: the row that fails first is the one to fix, and every row after it is downstream evidence.
An inverted mask polarity lands in the same place, because inverting a causal relation blocks query 0 from every key:
masked_scores (2, 3, 5, 5) [B,Nh,Tq,Tk] allowed 10/25; no fully blocked query False
^^ FIRST DIVERGENCE: no fully blocked query False
first divergence: masked_scores
The ledger covers Levels 1 through 3, cheaply, at every stage, with the axis names written down beside the numbers. It does not reach Level 4, which is why the reference comparison and the causal intervention are separate tools rather than more ledger rows.
One backward pass
Attention that only works in inference is not finished. A single smoke test establishes that the module supports training at all:
x = torch.randn(2, 8, 32, requires_grad=True)
loss = attn(x, is_causal=True).square().mean()
loss.backward()
loss: 0.033648 x.grad finite: True norm 0.010129
q_proj.weight grad present finite True norm 0.008869
k_proj.weight grad present finite True norm 0.008957
v_proj.weight grad present finite True norm 0.087069
out_proj.weight grad present finite True norm 0.064535
Four parameters receive finite gradients, and a finite gradient reaches the input. Chapter 3 gave the machinery for interpreting that evidence: every expected parameter has a path to the loss and nothing in this smoke test detached the graph. The particular gradient norms printed above are observations from one initialization, not a general ranking of Q, K, V and output-projection gradient sizes.
That is where this chapter stops on training. Whether those gradients have useful magnitudes, whether the optimizer owns the parameters, and whether the loss actually falls are Chapter 11’s subject.
Using AI on attention code
An assistant reading attention code sees operations that are individually correct, because they are. reshape is correct. softmax is correct. The mask has a legal shape. What the code does not contain is what each axis is supposed to mean, and that has to be established before any fix is proposed.
Here is a PyTorch attention implementation.
Do not rewrite it and do not propose fixes yet.
For every intermediate tensor:
1. Write its exact shape symbolically.
2. Give every axis a semantic name from: B, Nh, Tq, Tk, E, Dh, Dv.
3. State what ONE scalar or ONE vector at that stage represents.
4. State the invariant that should become true after that operation.
In particular:
- prove that split_heads followed by merge_heads is a value-preserving
round trip on a deterministic tensor, before analyzing anything else;
- derive the shape of Q @ K.transpose(-2, -1) and say which two axes moved;
- identify which axis holds the candidate keys, and verify that softmax
normalizes over that axis;
- state the boolean-mask convention for the exact PyTorch function being
called, and quote the documentation for that function;
- verify that forbidden positions receive zero probability, and check
whether any query has every key masked;
- say whether 1/sqrt(Dh) scaling is present and what it stabilizes;
- derive the shape of weights @ V and say which axis is contracted;
- verify that the merged result restores the intended external layout.
Identify the FIRST stage whose shape, axis meaning, value arrangement or
invariant disagrees with the intended computation.
Do not suggest fixes downstream of that first divergence.
The clause that earns its place is the one about quoting documentation for the exact function. Mask polarity is precisely the kind of detail that can be conflated across APIs when an assistant reasons from memory; the function-specific contract is stronger evidence than recollection.
That prompt reaches Levels 1 through 3. For Level 4, ask for an experiment rather than an opinion:
This attention module is supposed to be causal. Do not inspect the mask.
Write an experiment that changes only input positions after index i, reruns
the module in eval mode with dropout disabled, and reports the maximum
absolute change in the output at positions 0..i and at positions i+1 onward.
Before running it, tell me what each of those two numbers would have to be
for the causal contract to hold, and what it would mean if the second one
were also zero.
The last clause is the one that matters. An experiment where both numbers are zero has not proven causality; it has proven that the intervention did nothing.
Ask AI to name the query and key axes before asking it to fix the attention.
The attention debugging sequence
1. State the external contract: [B, T, E], and which axis is which. L1
2. Verify E is deliberately split: E == Nh * Dh. L1
3. Prove split_heads -> merge_heads is an exact round trip on a
deterministic tensor. Compare values, not shapes. L2
4. Label Q [..., Tq, Dh] and K [..., Tk, Dh], then derive
Q @ K.transpose(-2, -1) -> [..., Tq, Tk] and say which axes moved. L1/L2
5. Inspect score scale before softmax. Is 1/sqrt(Dh) present? L4
6. State the mask semantics for the exact API being called.
Does True mean allowed or blocked, for this function? L2
7. Apply the mask, prove the intended relations are excluded, and
check that no query has every key blocked. L3
8. Softmax over Tk. Verify every valid query row sums to one. L3
9. Derive weights [..., Tq, Tk] @ V [..., Tk, Dv] -> [..., Tq, Dv]. L1
10. Merge heads and verify the external output layout. L1/L2
11. For causal attention, run the future-token intervention test. L4
12. Compare against F.scaled_dot_product_attention under identical
Q, K, V, scale, equivalent mask and zero dropout. L4
13. Fix the FIRST failed invariant. Rerun. Do not fix two things at once.
Steps 1 through 4 catch the opening failure. Steps 6 through 8 catch the mask anchors. Steps 11 and 12 catch what nothing before them can.
Exercises
Same shape, wrong heads. Reproduce the opening split with a deterministic
arangetensor. Show the shapes, dtypes and element multisets are equal and the tensors are not. Run both through the same attention and report the maximum output difference, then write down which positions ended up in “head 0” under the bad split.Split/merge round trip. Implement both functions and property-test
merge_heads(split_heads(x, Nh)) == xacross several[B, T, E]shapes and everyNhdividingE. Then break the merge by reshaping directly from[B, Nh, T, Dh]and confirm the test catches it.The scale sweep. For
Dhin{8, 32, 128, 512}, measure the score standard deviation, mean maximum softmax probability and mean softmax entropy, with and without1/sqrt(Dh). Compare the unscaled standard deviation withsqrt(Dh). State the assumptions under which the relationship holds.Wrong softmax axis. Normalize scores over
dim=-2instead ofdim=-1. Show which sums become one, explain what relationship each normalization describes, and explain why no shape assertion can tell them apart.Mask polarity. Build one causal allow relation and send it through
manual_attention,F.scaled_dot_product_attentionandnn.MultiheadAttention, converting the polarity explicitly for each. Then send the wrong polarity to each and record which produces wrong finite numbers, which producesNaN, and which query rows are affected.Causal intervention. Take a causal module in eval mode. Change only positions after
iand prove the outputs at0..ido not move, while the outputs afterido. Repeat with causality disabled and with the mask polarity inverted, and record the three sets of numbers.Cross-attention. With
Tq = 7,Tk = 11,Dh = 16andDv = 5, derive every shape before executing anything. Then show that a square causal mask does not broadcast, build a correct key-padding mask, and verify that padded keys receive zero weight.Sort the failures by level. Break
manual_attentionone way at a time: remove the scale, flip the softmax axis, skip the head transpose, invert the mask. For each, record which of the four levels first detects it, and which checks pass all the way through.
Next: it all runs, and it still does not learn
Attention is no longer a black box with a shape contract. It is a short sequence of relationships between represented vectors, and every stage can be named, derived and checked:
x [B, T, E] T represented vectors per example
Q, K, V [B, T, E] three learned projections, nothing new yet
split [B, Nh, T, Dh] feature axis partitioned, head axis moved outward
scores [B, Nh, Tq, Tk] query i against key j, scaled by 1/sqrt(Dh)
masked [B, Nh, Tq, Tk] forbidden relations excluded before normalization
weights [B, Nh, Tq, Tk] each valid query row a distribution over keys
context [B, Nh, Tq, Dv] a convex combination of value vectors per query
merged [B, Tq, E] head order restored, round trip provable
output [B, Tq, E] external contract preserved
Five deliberate failures ran through that pipeline, and together they show why the levels have to be layered rather than collapsed into one check:
Shape tells you how many relationships exist. It does not tell you whether they are the right relationships. Name the query and key axes, derive each transformation, verify the invariant that stage introduces, and stop at the first place where the relationships stop meaning what you intended.
That is real progress, and it is worth being precise about how much. We can now establish that the head values are arranged correctly, that the mask means what this API thinks it means, that causality holds under intervention, that every query row is a proper distribution over allowed keys, that the implementation agrees with PyTorch’s own, that the external layout is preserved, and that finite gradients reach every parameter.
None of that establishes that the model will learn anything. A mechanically correct attention block sits inside a system with data, targets, a loss, an optimizer, a learning rate and a schedule, and any one of them can be broken in a way that produces no exception at all. The attention-specific four-level ladder is no longer sufficient once that happens. The next chapter develops a different set of experiments for the training system itself — beginning from the harder case where the entire program runs, the loss is finite, the optimizer steps, and the model is still useless.
Into a model that runs and does not learn.