Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Why apply softmax along dim=-1 in attention, and what if along another dimension?
What you need to know
Reading the axes
scores: (batch, heads, T_q, T_k) row i = one query, one score per key column j = one key, scored by every querySoftmax over dim=-1 answers "for query i, how should its attention be shared among keys?". That is what weights @ V needs: for each query, weights over keys that sum to 1.
What goes wrong along dim=-2
- No longer an average. Each column sums to 1, but rows do not. A query's output is a sum of value vectors with an arbitrary total weight, which drifts with sequence length.
- Causality breaks. Column
jholds scores from all queriesi ≥ j. Normalising down that column divides token 0's weight on key 0 by a sum that includes the scores of tokens 1, 2, 3 and so on. Token 0's output now depends on future tokens.
This test checks which inputs affect the output at token 0:
1import torch2torch.manual_seed(0)3T, hd = 5, 84x = torch.randn(T, hd, requires_grad=True) # token inputs5future = torch.triu(torch.ones(T, T, dtype=torch.bool), 1)67for dim in (-1, -2):8 s = (x @ x.T).masked_fill(future, float("-inf"))9 out = s.softmax(dim=dim) @ x10 out[0].sum().backward() # which inputs affect token 0?11 print(dim, x.grad.abs().sum(-1).gt(0).tolist())12 x.grad = None13# -1 [True, False, False, False, False]14# -2 [True, True, True, True, True]With dim=-1, token 0 depends only on itself. With dim=-2, it depends on all five tokens — including four it must never see.
Why the training loss "looks great"
During training, the model gets the whole sequence at once. With the leak, position t can pick up information about token t+1, the very token it must predict, so loss falls far faster than it should. At generation time the future does not exist yet, so the learned shortcut fails and the output is nonsense.
A real-life example
A code-completion team writes a custom attention kernel to support a new tokenizer. After a day of training, validation loss is far below their previous best, and someone suggests announcing a breakthrough. A senior engineer runs one quick generation instead: the model repeats a few tokens in a loop.
She runs the dependency test above on their kernel and finds the output at position 0 has gradient from every later position. The softmax was applied over the wrong axis after a transpose. They add that test to CI, alongside a check that every attention row sums to 1, and retrain. Validation loss returns to a normal level — and the model completes code again.
Follow-up questions to expect
- "How would you catch this bug quickly?" — Check that each row of the attention matrix sums to 1, and that perturbing a later token does not change an earlier token's output.
- "Is
dim=-1right for cross-attention too?" — Yes. Each decoder query still needs a distribution over the encoder's keys, which are on the last axis. - "Why is 'loss too good' a warning sign?" — Language modelling has a floor set by real uncertainty in text; a sudden large drop usually means the model can see the answer.