Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Describe how causal mask is combined with attention scores and why order matters?
What you need to know
The correct order
1. scores = Q Kᵀ / sqrt(d_k) (T × T)2. scores[i, j] = -inf for every j greater than i3. weights = softmax(scores, over keys)4. out = weights · VThe mask is upper-triangular: row i may see columns 0..i. It is built once, for example with torch.triu(ones, diagonal=1), and sliced to the current length.
Showing the difference
1import torch2torch.manual_seed(0)3T, hd = 4, 84q, k = torch.randn(T, hd), torch.randn(T, hd)5scores = q @ k.T / hd**0.56future = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)78right = scores.masked_fill(future, float("-inf")).softmax(-1)9wrong = scores.softmax(-1).masked_fill(future, 0.0)10print(right.sum(-1)) # tensor([1., 1., 1., 1.])11print(wrong.sum(-1)) # tensor([0.1211, 0.5228, 0.5401, 1.0000])With the wrong order, token 0 keeps only 12% of its weight. Its output vector shrinks, and by how much depends on the scores of tokens 1–3 — tokens it must not see. The last row is unaffected because nothing is in its future, which is why the bug can hide in a quick test that only looks at the last position.
Why add minus infinity, not multiply by zero
exp(-inf) = 0, so masked positions get exactly zero weight and the rest renormalise correctly. Setting a score to 0 would give it weightexp(0) = 1, which is not masking at all.- Masks are additive, so a causal mask, a padding mask and a position bias such as ALiBi combine by simple addition before one softmax.
Practical details
- Many implementations use
torch.finfo(dtype).mininstead of-inf. If a row is fully masked (for example, a padding query), all-infgivesNaNafter softmax, while a large finite value gives a harmless uniform row. - Scale before masking. With
-infthe order of scale and mask does not matter, but with a finite sentinel, dividing it bysqrt(d_k)makes it less negative. - In practice, use
F.scaled_dot_product_attention(q, k, v, is_causal=True). The fused kernels apply the mask inside the tiles and skip fully masked blocks entirely, saving about half the work.
A real-life example
A code-completion team writes its own attention function to try a new position scheme. They apply the mask after the softmax. Training loss looks normal and the model can still complete code, so no one notices for a week.
Then they compare against the reference scaled_dot_product_attention on the same checkpoint and find outputs differ by up to 40% in early positions. Early tokens in every file have been attenuated during training, and their scale leaked information about later tokens. They fix the order, add a unit test that compares every row, not just the last, against the reference, and retrain.
Follow-up questions to expect
- "Why does the last row look correct even with the bug?" — It has no future positions, so nothing is masked and the row already sums to 1.
- "What happens with a fully masked row?" — With
-infeverywhere, softmax divides zero by zero and returns NaN; use a finite minimum or skip those rows. - "Where does the padding mask go?" — In the same place, added before the softmax, so padded keys get zero weight.