Causal masking
Block every token from reading anything that comes after it, so a model can be trained on every position of a sentence at once without cheating.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A causal mask stops every word from seeing the words that come after it.
Think about how you were taught to read aloud in school. You put a ruler or a folded sheet under the line you were on, covering everything below. You could see what you had already read, and nothing more.
Slide the ruler down one line and a little more of the page appears. What you had already read never changes. Only the amount you can see grows.
That sliding ruler is the causal mask. "Causal" here means the past can affect the future, never the reverse.
Why a language model needs it
A language model is trained on one task: given the words so far, guess the next word.
Now think about how you would train it. You have a sentence of twenty words. You could show it the first word and ask for the second. Then the first two, and ask for the third. And so on. Nineteen separate passes for one sentence. Painfully slow.
The alternative is to show it the whole sentence at once. Every position guesses its own next word, all in parallel. Nineteen answers from one pass.
There is an obvious problem with that. Suppose word five can see word six. Then the question "what comes after word five?" has its answer sitting right there. The model would learn to copy rather than predict. Training loss would drop to almost nothing. The model would be useless the moment you asked for something new.
The causal mask is what makes the fast version safe. Every position may look backwards and at itself. Nothing may look forward.
How the blocking is done
You cannot delete the future words, because they are needed for their own predictions. So instead the scores against future words are pushed to negative infinity before the shares are worked out.
The step that turns scores into shares gives a share of exactly zero to a score of negative infinity. The future words are still sitting in memory. They receive no weight at all.
the cat sat down (what is being looked at)
the . X X X
cat . . X X
sat . . . X
down . . . .
. = allowed X = blocked
(rows are the word doing the looking)The allowed region is a triangle. The first word sees only itself. The last word sees everything. Every row in between sees a little more than the row above.
The property that makes generating text cheap
Here is the consequence that matters most, and it is worth checking rather than believing.
Run the model on three words. Note what comes out for each of them. Now run it on those same three words plus a fourth. The results for the first three are exactly identical, down to the last digit.
That has to be true. The first three words could not see the fourth, so adding it cannot have changed them.
This is why writing text is fast. When a model produces a new word, it does not redo the work for everything before it. It keeps the earlier keys and values in a store called the KV cache. That is a saved copy of what each earlier word advertised and offered. Only the new word gets computed.
Take away the mask and that guarantee is gone. Every new word would change every earlier word, and each step would cost as much as starting over.
What this does not do
The mask does not make the model correct. It does not stop it inventing facts. It stops one specific form of cheating during training.
It is also the difference between the two big families of language model. Masked, and the model writes. Unmasked, and the model reads and understands but cannot write left to right. That split is covered in decoder-only vs encoder-decoder.
Remember this
- The mask blocks every word from reading anything to its right.
- It lets one pass train every position at once, without any position seeing its own answer.
- Because the past never changes, earlier work can be cached, which is what makes text generation fast.
What to learn next
- Cross-attention — attention between two different sequences, where masking works differently.
- How LLMs work — the mask in the context of the whole generation loop.
- Context window — why the cache, not the weights, sets the practical length limit.
Developer — Code and libraries.
Setup
pip install numpyThe mask, and the property it guarantees
import numpy as np
rng = np.random.default_rng(3)
np.set_printoptions(precision=3, suppress=True)
def softmax(x, axis=-1):
x = x - x.max(axis=axis, keepdims=True)
e = np.exp(x)
return e / e.sum(axis=axis, keepdims=True)
def attention(X, Wq, Wk, Wv, causal):
Q, K, V = X @ Wq, X @ Wk, X @ Wv
logits = Q @ K.T / np.sqrt(Q.shape[-1])
if causal:
T = X.shape[0]
future = np.triu(np.ones((T, T), dtype=bool), k=1) # True above the diagonal
logits = np.where(future, -np.inf, logits) # -inf becomes 0 after softmax
A = softmax(logits)
return A, A @ V
d = 4
Wq, Wk, Wv = (rng.normal(size=(d, d)) for _ in range(3))
words = ["mera", "naam", "Pranay", "hai"]
X = rng.normal(size=(4, d))
A_open, _ = attention(X, Wq, Wk, Wv, causal=False)
A_causal, out4 = attention(X, Wq, Wk, Wv, causal=True)
print("without a mask - every token sees every token:")
print(" " + "".join(f"{w:>9s}" for w in words))
for w, row in zip(words, A_open):
print(f"{w:>9s} " + "".join(f"{v:9.3f}" for v in row))
print("\nwith a causal mask - nobody sees the future:")
print(" " + "".join(f"{w:>9s}" for w in words))
for w, row in zip(words, A_causal):
print(f"{w:>9s} " + "".join(f"{v:9.3f}" for v in row))
print("\nthe mask itself (True = blocked):")
print(np.triu(np.ones((4, 4), dtype=bool), k=1))
# The property that makes training a language model possible at all.
_, out3 = attention(X[:3], Wq, Wk, Wv, causal=True)
print("\nrun on 3 tokens, then on 4. do the first 3 outputs change?")
print(" max difference:", np.abs(out4[:3] - out3).max())
_, open3 = attention(X[:3], Wq, Wk, Wv, causal=False)
_, open4 = attention(X, Wq, Wk, Wv, causal=False)
print(" same test without the mask:", np.abs(open4[:3] - open3).max())without a mask - every token sees every token:
mera naam Pranay hai
mera 0.030 0.043 0.299 0.628
naam 0.042 0.092 0.841 0.026
Pranay 0.694 0.306 0.000 0.000
hai 0.001 0.001 0.007 0.991
with a causal mask - nobody sees the future:
mera naam Pranay hai
mera 1.000 0.000 0.000 0.000
naam 0.312 0.688 0.000 0.000
Pranay 0.694 0.306 0.000 0.000
hai 0.001 0.001 0.007 0.991
the mask itself (True = blocked):
[[False True True True]
[False False True True]
[False False False True]
[False False False False]]
run on 3 tokens, then on 4. do the first 3 outputs change?
max difference: 0.0
same test without the mask: 4.6969765681292355Read the last two lines first
0.0 and 4.697. That contrast is the whole lesson.
With the mask, extending the sequence changed the earlier outputs by exactly zero. Not approximately. Not to seven decimal places. Zero. Without the mask, the same tokens moved by 4.7, which for these vectors is an enormous change.
That exact-zero property is what a KV cache relies on. If it were only approximately zero, incremental generation would drift away from a full forward pass. Nobody would trust it.
Make this your regression test. Whenever you touch masking, run the model on a prefix and on the full sequence. Assert the prefix outputs are identical.
Two details in the attention tables
The first row is 1.000 and three zeros. The first token can only see itself, so all of its weight lands on itself. Its output equals its own value vector, unchanged. The first token of a causal model is always a pass-through.
The Pranay row is identical in both tables. That is a coincidence in this seed, not a rule. That token's raw scores against the future were already very negative. Masking them changed almost nothing. Do not read a pattern into it.
The right way to do this in PyTorch
import torch.nn.functional as F
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)Written against PyTorch 2.5.1; unchanged in the current 2.13 documentation. Prefer is_causal=True over building a mask yourself:
- The fused kernel skips the blocked blocks entirely, roughly halving the attention work.
- No
TbyTmask tensor is allocated. - Passing both
attn_maskandis_causal=Trueraises an error. They are mutually exclusive.
For an explicit mask, PyTorch accepts a boolean tensor where True means "keep". It also accepts a float tensor added to the logits. These conventions are opposite to each other and mixing them up is a common bug.
Combining a causal mask with a padding mask
Batched sequences have different lengths, so short ones are padded. Padding must be masked too, and the two masks combine with a logical AND.
This creates a real trap. A padded row can end up with every key blocked. A softmax over an all-blocked row produces nan, which spreads through the whole batch. Two safe options. Use a large finite negative number instead of negative infinity. Or drop padded positions from the loss, so their outputs never matter. See softmax overflow inside attention.
Common mistakes
k=0 instead of k=1 in np.triu. With k=0 the diagonal is blocked too, so no token can see itself. Loss stays stubbornly high and nothing about the error message points at the mask.
Masking after the softmax. Zeroing weights afterwards leaves rows that no longer sum to one. The output is then scaled down by an amount that varies per row. Mask the logits, always.
Building the mask fresh in every forward pass. For a fixed maximum length, build it once and register it with register_buffer. It then moves with the model between devices. See buffers vs parameters.
Assuming an encoder needs a causal mask. A model that classifies or embeds text should read the whole sentence. Masking a classifier cripples it for no benefit.
Try it yourself
Add a sliding-window mask on top of the causal one. Each token then sees only itself and the two before it. Build it with np.tril(np.ones((T, T)), k=0) - np.tril(np.ones((T, T)), k=-3). Then check the shape of the allowed band. Then re-run the prefix test: the exact-zero property still holds, and the cache gets smaller.
What to learn next
- Cross-attention — attention between two different sequences, where masking works differently.
- How LLMs work — the mask in the context of the whole generation loop.
- Context window — why the cache, not the weights, sets the practical length limit.
Researcher — Mathematics and papers.
The mask as an additive term
$$ A = \operatorname{softmax}!\left( \frac{QK^\top}{\sqrt{d_k}} + M \right), \qquad M_{ij} = \begin{cases} 0 & j \le i \ -\infty & j > i \end{cases} $$
Since $\exp(-\infty) = 0$, each row $i$ normalises over the $i+1$ allowed positions. This makes the model autoregressive: the joint distribution factorises as
$$ p(x_1, \dots, x_T) = \prod_{t=1}^{T} p(x_t \mid x_{<t}) $$
and a single forward pass evaluates all $T$ conditionals. That is the entire efficiency argument for the mask, and it is why decoder-only pretraining scales.
Prefix invariance, stated precisely
Let $f_\theta$ be a causally masked transformer. For any sequence $x_{1:T}$ and any $t \le T$:
$$ f_\theta(x_{1:T}){1:t} = f\theta(x_{1:t}) $$
This is an exact equality on real arithmetic, and it is the correctness condition for KV caching. In floating point it holds bitwise only when the kernel's reduction order is identical between the two calls. Fused attention kernels choose tile sizes based on sequence length. A prefill of length $T$ and an incremental decode step can differ in the last bits. Divergence between batched-prefill and incremental-decode paths is a real source of serving bugs. It is also under-reported.
Cost
A causal kernel needs only the lower triangle: $T(T+1)/2$ score entries rather than $T^2$. FlashAttention with is_causal skips entire blocks above the diagonal, giving close to the theoretical factor of two. A naive implementation that materialises the full matrix and then adds $M$ pays the full $T^2$ and gains nothing.
During autoregressive decoding, step $t$ costs $O(t d)$ per layer against the cache. Generating $n$ tokens therefore costs $O(n^2 d)$ per layer. The KV cache itself is $2 L g d_k$ bytes per token. Here $L$ is depth and $g$ the number of key-value groups. That constraint motivates grouped-query attention and every KV compression method.
Variants of the mask
- Prefix language modelling applies bidirectional attention over a prefix and causal attention over the continuation. See UniLM (Dong et al., 2019) and T5 (Raffel et al., 2020). The mask is block-structured rather than triangular. That lets a single decoder stack behave like an encoder over the prompt.
- Sliding-window attention intersects the causal mask with a band of width $w$, giving $O(Tw)$ cost. Stacking $L$ such layers yields an effective receptive field of $Lw$. Used in Longformer (Beltagy et al., 2020) and Mistral 7B (arXiv:2310.06825).
- Document masking blocks attention across document boundaries inside a packed training batch. Omitting it lets a model attend across unrelated documents that share a batch row. That measurably degrades long-context behaviour.
- Speculative and tree decoding need non-triangular masks so several candidate continuations can be verified in one pass. This is the main reason a serving stack keeps a general mask path alongside the fast causal one.
An empirical caution about position zero
Xiao et al. (2024), arXiv:2309.17453, show a consistent effect. Causal models place disproportionate attention on the first few tokens, regardless of content. The proposed explanation is structural. Every row's softmax must sum to one, and position 0 is the only key visible to every query. It becomes the default destination for surplus attention mass. Evicting those tokens from a rolling KV cache degrades quality sharply; keeping four of them restores it. Any cache-eviction policy has to account for this.
Papers
- Vaswani et al., Attention Is All You Need, 2017 — arxiv.org/abs/1706.03762
- Radford et al., Improving Language Understanding by Generative Pre-Training, 2018
- Dong et al., Unified Language Model Pre-training, 2019 — arxiv.org/abs/1905.03197
- Beltagy et al., Longformer, 2020 — arxiv.org/abs/2004.05150
- Xiao et al., Efficient Streaming Language Models with Attention Sinks, 2024 — arxiv.org/abs/2309.17453
What to learn next
- Cross-attention — attention between two different sequences, where masking works differently.
- How LLMs work — the mask in the context of the whole generation loop.
- Context window — why the cache, not the weights, sets the practical length limit.