Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Define causal masking and explain when and why it is applied during attention?
What you need to know
Teacher forcing makes training fast: a 2,048-token sequence gives 2,047 predictions from one forward pass. But all 2,048 tokens are in the input, so each position could see the tokens it is supposed to predict. The causal mask removes that shortcut.
Worked by hand
Take the scaled scores for 3 tokens from the previous section's example:
scaled scores mask (x = future) weights after softmax[1.0 0.0 0.5] [. x x] [1.000 0 0 ][0.0 1.0 0.5] -> [. . x] -> [0.269 0.731 0 ][1.0 0.5 1.0] [. . .] [0.384 0.233 0.384]Token 1 can only see itself, so its weight is 1.0. Token 2 splits its weight between tokens 1 and 2. Token 3 sees everything. Each row still sums to 1.
1import torch, torch.nn.functional as F23T = 34scores = torch.tensor([[1.0, 0.0, 0.5],5 [0.0, 1.0, 0.5],6 [1.0, 0.5, 1.0]]) # already scaled Q Kᵀ / sqrt(d_k)7future = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)8weights = scores.masked_fill(future, float("-inf")).softmax(dim=-1)9print(weights)1011# The built-in does the same when is_causal=True12q = k = v = torch.randn(1, 1, T, 4)13manual = ((q @ k.transpose(-2, -1)) / 2).masked_fill(future, float("-inf")).softmax(-1) @ v14print(torch.allclose(manual, F.scaled_dot_product_attention(q, k, v, is_causal=True), atol=1e-6))15# Truetorch.triu(..., diagonal=1) marks the cells strictly above the diagonal — the future. In real code, pass is_causal=True to scaled_dot_product_attention so the fused kernel can skip those cells entirely.
When it applies
| Phase | Mask needed? | Why |
|---|---|---|
| Training (decoder) | Yes | All positions are processed together |
| Prefill (reading the prompt) | Yes | Same reason: many positions at once |
| Decode with KV cache | Implicit | The one new query only sees cached past keys |
| Encoder self-attention | No | Bidirectional by design; only padding is masked |
| Cross-attention | No | The whole source is known in advance |
The mask also makes training match inference: at generation time future tokens do not exist, so the model must never learn to rely on them.
A real-life example
A team trains a small code-completion model on Python files. After one hour, training loss is almost zero, but at inference the model produces random tokens. The cause was a refactor that replaced is_causal=True with a padding-only mask. Each position could see the next token, so the model learned to copy it — a task that is trivial in training and impossible in generation.
A quick unit test catches this: change the last token of an input and check that the logits at every earlier position stay exactly the same. If they change, information is leaking from the future.
Follow-up questions to expect
- "Why is the mask applied before softmax?" — So the masked positions get exactly zero weight and the visible ones still sum to 1.
- "Is the mask a learned parameter?" — No. It is a fixed pattern; it only depends on sequence length.
- "Does a causal mask save compute?" — Fused kernels skip the fully masked tiles, saving roughly half the attention work.