Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Define causal masking and explain when and why it is applied during attention?


The same scores with a causal mask1.00··0.270.73·0.380.230.38token 1token 2token 3token 1token 2token 3Dotted cells were set to minus infinity before softmax.
Row two renormalises to 0.27 and 0.73 over the two tokens it may see — the mask removes the future without breaking the sum to one.

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:

Text
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.

Python
import torch, torch.nn.functional as FT = 3scores = torch.tensor([[1.0, 0.0, 0.5],                       [0.0, 1.0, 0.5],                       [1.0, 0.5, 1.0]])               # already scaled Q Kᵀ / sqrt(d_k)future = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)weights = scores.masked_fill(future, float("-inf")).softmax(dim=-1)print(weights)# The built-in does the same when is_causal=Trueq = k = v = torch.randn(1, 1, T, 4)manual = ((q @ k.transpose(-2, -1)) / 2).masked_fill(future, float("-inf")).softmax(-1) @ vprint(torch.allclose(manual, F.scaled_dot_product_attention(q, k, v, is_causal=True), atol=1e-6))# True

torch.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

PhaseMask needed?Why
Training (decoder)YesAll positions are processed together
Prefill (reading the prompt)YesSame reason: many positions at once
Decode with KV cacheImplicitThe one new query only sees cached past keys
Encoder self-attentionNoBidirectional by design; only padding is masked
Cross-attentionNoThe 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.