Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Why apply softmax along dim=-1 in attention, and what if along another dimension?


Which inputs can change token 0's output✓····✓✓✓✓✓t0t1t2t3t4softmax dim=-1softmax dim=-2Measured with gradients on a 5-token causal toy.
Normalising over queries lets every future token reach back into token 0, which is why the training loss looks too good to be true.

What you need to know

Reading the axes

Text
scores: (batch, heads, T_q, T_k)         row i    = one query, one score per key         column j = one key, scored by every query

Softmax 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 j holds scores from all queries i ≥ 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:

Python
import torchtorch.manual_seed(0)T, hd = 5, 8x = torch.randn(T, hd, requires_grad=True)       # token inputsfuture = torch.triu(torch.ones(T, T, dtype=torch.bool), 1)for dim in (-1, -2):    s = (x @ x.T).masked_fill(future, float("-inf"))    out = s.softmax(dim=dim) @ x    out[0].sum().backward()                      # which inputs affect token 0?    print(dim, x.grad.abs().sum(-1).gt(0).tolist())    x.grad = None# -1 [True, False, False, False, False]# -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=-1 right 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.