Causal Masking
Goal
Construct and test a causal mask that makes future attention probability exactly zero for next-token prediction.
For a decoder predicting left to right, query position t may use positions 0..t and must not use t+1....
For length 4, query position 2 may attend to keys 0, 1, and 2. Key 3 is future information. The standard numerical move is to replace illegal logits with negative infinity before softmax; softmax then assigns those positions probability zero.
Why masking must happen before normalization
Take unmasked logits [2, 1, 3] for a query that is allowed to use only the first two positions. If you softmax all three first, the illegal future logit 3 steals probability mass. Simply setting its probability to zero afterward leaves the legal weights summing to less than one.
The correct procedure is:
raw logits: [2, 1, 3]
mask future: [2, 1, -∞]
softmax legal: [~0.73, ~0.27, 0]
Now all normalized mass is redistributed only among legal positions.
This is more than an implementation preference. If an earlier token can use a later target token during training, the task becomes easier for the wrong reason. Loss can look excellent while the model is learning from information it will not have during generation.
That is why prefix invariance is such a strong end-to-end check: two sequences that share the same prefix but differ only later should produce the same earlier outputs. A triangular-looking heatmap is suggestive; unchanged prefix outputs are behavioral evidence that the future did not leak through the stack.
See the triangular legality pattern
For a four-token decoder sequence, the legality pattern is:
query 0: can use 0
query 1: can use 0,1
query 2: can use 0,1,2
query 3: can use 0,1,2,3
Written as a matrix:
1 0 0 0
1 1 0 0
1 1 1 0
1 1 1 1
Each row is a query position; each column is a key/value position.
The zeros are not “low attention preferences.” They are illegal information paths for next-token training.
This distinction matters. A low learned score can change with training. A causal mask encodes a task rule that must remain enforced regardless of learned parameters.
Later, when generation runs one token at a time, there is no future token to leak. The mask is crucial during parallel training because all target tokens are physically present in the training tensor.
Predict
Turn causality into an executable invariant
Run the starter once and observe non-zero future attention. Complete the TODO in causal_logits so positions after the query index receive NEG_INF before softmax. Run again and confirm future attention mass is exactly zero. Then deliberately unmask one future entry and observe the invariant fail.
Loading lab…
The strongest simple check is not “the heatmap looks triangular.” It is numerical: future attention mass = 0. Another end-to-end check is prefix invariance: changing only future tokens must not alter outputs for an earlier shared prefix.
Quick Check
Explain it back
Explain why masking after softmax and simply zeroing forbidden weights without renormalizing changes the meaning of the remaining mixture.
Key Takeaways
- Causal masking is a correctness contract for decoder attention.
- Future logits must be blocked before softmax.
- Future attention mass should be exactly zero.
- Prefix invariance is a powerful end-to-end causal test.
Next Lesson
Complete the checkpoint now. Then split the feature channels into multiple attention heads while preserving the same causal rule.
References
- Vaswani et al., Attention Is All You Need.
Completion is stored locally on this device.