Every self-attention layer computes a score between each query and each key, turns each row of scores into weights with a softmax, and mixes the values. Causal and bidirectional attention run exactly this computation with the same weights. The only difference is which scores are allowed to survive the softmax. In bidirectional attention every token can read every other token. In causal attention token i can read only tokens 0 to i. That one triangle of masked scores decides what the model can be trained to do, how it is served, and which bugs it is exposed to.
This page explains the mask from first principles, shows why it follows from the training objective rather than from taste, works through the inference consequences (the KV cache only exists because of causality), and covers hybrids such as prefix LMs. It ends with the PyTorch code, the alignment trap that silently breaks decoding, a test that catches future leakage, and a checklist.
The mask is the whole difference
Write the attention scores as S = QKT / sqrt(d), a T by T matrix for a sequence of T tokens. Row i holds query i's scores against every key. Masking means adding minus infinity to the forbidden entries before the softmax, so exp(minus infinity) = 0 and those keys receive exactly zero weight. For causal attention the forbidden entries are those with column j greater than row i: the strict upper triangle.
A worked example makes it concrete. Take the three tokens "the cat sat" and suppose one head's raw scores for the query at "cat" are 1.0 for "the", 2.0 for "cat" and 3.0 for "sat". Bidirectionally the softmax gives weights of about 0.09, 0.24 and 0.67, so most of the output comes from "sat", a token to the right. Causally the score for "sat" becomes minus infinity, and the remaining two renormalise to about 0.27 and 0.73. The representation of "cat" now depends only on "the cat".
Because the mask is applied inside every layer, the restriction composes. In a causal stack, position i's representation at layer 12 is a function of tokens 0 to i only, however many layers deep you go. In a bidirectional stack, every position's representation depends on the whole sequence from the first layer onward.
Why the objective dictates the mask
A causal language model is trained to predict token i+1 from tokens 0 to i, at every position at once. That is only a fair test if position i cannot read token i+1. With the causal mask, one forward pass over a T-token sequence yields T valid next-token predictions, each computed as if the future did not exist. This is teacher forcing: the true prefix is fed in, all predictions are scored in parallel, and generation later reproduces exactly that setting one token at a time.
Remove the mask from a next-token model and training loss collapses towards zero, because every position can copy its answer from the next column. The model has learned nothing usable, and at generation time, where the future really is absent, it fails. A falling loss with no matching improvement in samples is the classic symptom of this leak.
A bidirectional encoder needs a different objective, because reading the whole input makes next-token prediction trivial. Masked language modelling hides some input tokens and asks the model to reconstruct them from both sides. BERT selected 15 percent of tokens; of those, 80 percent were replaced by a [MASK] token, 10 percent by a random token and 10 percent left unchanged, so that the model could not rely on seeing [MASK] at fine-tuning time. The price is sample efficiency: only the selected positions contribute to the loss, whereas a causal model gets a training signal from every position.
What each mask buys and costs
| Property | Bidirectional | Causal |
|---|---|---|
| Context per token | Whole sequence, both sides | Only earlier tokens |
| Natural objective | Masked or denoising reconstruction | Next-token prediction |
| Loss positions per sequence | The selected subset | All T positions |
| Generation | Not directly; needs an iterative or separate decoder | Native, one token per step |
| Appending a token | Every representation can change, so recompute all | Earlier representations are fixed, so cache K and V |
| Typical uses | Classification, tagging, retrieval embeddings, rerankers | Chat, code completion, any open-ended generation |
| Attention FLOPs | Full T by T | About half, if the kernel skips masked blocks |
The last row deserves care. The mask removes half of the score matrix, but a naive implementation still computes all of it and then discards half. Fused kernels such as FlashAttention skip tiles that lie entirely above the diagonal, so causal attention costs roughly half the attention FLOPs in practice. The projections and MLP are unaffected, so the end-to-end saving is smaller.
Representation quality is a separate question. For a fixed size, a bidirectional encoder usually gives better token and sentence representations for understanding tasks, because every token sees its right context. Large causal models close much of that gap through scale, and many embedding models are now built from decoders, sometimes by fine-tuning with the causal mask removed.
Inference: why the KV cache needs causality
In a causal model, appending token T cannot change the keys and values of tokens 0 to T-1, because none of them could ever read token T. Their K and V vectors at every layer are final the moment they are computed. Decoding therefore stores them in a KV cache and, at each step, computes Q, K and V for the new token only, attending from one query to all cached keys. The cost per new token grows with context length rather than with its square, and prefill (the prompt) runs once.
A bidirectional model has no such property. Append a token and every earlier token's representation may change at every layer, so the whole sequence must be recomputed. That is why encoders are used where the input is complete before inference starts: classify, embed or rerank a finished text in one pass. For those jobs the encoder is cheap, since there is no decode loop at all.
Encoder-decoder models combine the two. The encoder reads the source bidirectionally once. The decoder uses causal self-attention over its own outputs plus unmasked cross-attention to the encoder states. The decoder's self-attention keys are cached as usual, and the cross-attention keys and values are computed once per request because the source never changes.
The masks in PyTorch
PyTorch's scaled_dot_product_attention accepts either is_causal=True or an explicit attn_mask. Its documentation (checked against the 2.14 page on 2026-10-02) states three things that matter here: a boolean mask value of True means the position takes part in attention; an error is thrown if both attn_mask and is_causal are set; and for non-square shapes is_causal produces an upper-left aligned triangle. Note that nn.MultiheadAttention uses the opposite boolean convention, where True means masked out, so masks cannot be passed between the two APIs unchanged.
import torch
import torch.nn.functional as F
B, H, T, D = 2, 8, 6, 64
q = k = v = torch.randn(B, H, T, D)
# 1. Bidirectional: no mask at all.
enc = F.scaled_dot_product_attention(q, k, v)
# 2. Causal: let the kernel build the triangle (square case only).
dec = F.scaled_dot_product_attention(q, k, v, is_causal=True)
# 3. The same causal mask written out. In SDPA a boolean True means "may attend".
causal = torch.ones(T, T, dtype=torch.bool).tril()
dec2 = F.scaled_dot_product_attention(q, k, v, attn_mask=causal)
assert torch.allclose(dec, dec2, atol=1e-5)
# 4. Prefix LM: the first P positions see each other in both directions.
P = 3
prefix = causal.clone()
prefix[:, :P] = True
plm = F.scaled_dot_product_attention(q, k, v, attn_mask=prefix)
# 5. Padding: combine with a key-padding mask by AND, never by replacing.
lengths = torch.tensor([6, 4])
key_ok = torch.arange(T)[None, :] < lengths[:, None] # (B, T)
mask = causal[None, None] & key_ok[:, None, None, :] # (B, 1, T, T)
out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)Two habits prevent most mask bugs. Combine masks by logical AND, so that padding never re-enables a future position. Build masks once per shape and assert their shape and dtype, because a mask broadcast along the wrong axis runs without error and produces plausible garbage.
The alignment trap in decoding and chunked prefill
During decoding the query block is shorter than the key block: two new queries attend to five keys, three cached plus the two just appended. The correct causal mask aligns the triangle to the bottom-right corner, so the last query sees every key and the first new query sees all but the last. An upper-left aligned triangle instead lets the first new query see only key 0. Nothing crashes, and outputs simply degrade.
Libraries disagree here. SDPA's is_causal is documented as upper-left for non-square shapes. The flash-attn README states that from version 2.1, when the query and key lengths differ, the causal mask is aligned to the bottom-right corner, while earlier versions used top-left. Code that swaps one backend for another, or that upgrades across that boundary, can change semantics without changing a line. Write the offset explicitly.
def decode_step(q_new, k_cache, v_cache):
# q_new: (B, H, n_new, D); caches already include the n_new new keys and values.
n_new, n_kv = q_new.shape[-2], k_cache.shape[-2]
offset = n_kv - n_new
# Bottom-right alignment: new query i may read keys 0 .. offset + i.
mask = torch.ones(n_new, n_kv, dtype=torch.bool).tril(diagonal=offset)
return F.scaled_dot_product_attention(q_new, k_cache, v_cache, attn_mask=mask)
# Wrong: is_causal=True here would apply the top-left triangle.
Hybrids between the two
- Prefix LM. Tokens in a prefix (an instruction, a document, an image's patches) attend to each other bidirectionally, and the continuation is causal. You get bidirectional context over the input and still generate. The cost is that prefix tokens cannot be scored by next-token loss, and adding to the prefix invalidates its cache.
- Encoder-decoder. Separate weights for the bidirectional reader and the causal writer, joined by cross-attention. A good fit when inputs and outputs differ in kind, such as translation or speech to text.
- Packed sequences. Training often concatenates several documents into one row. A plain causal mask lets document two attend to document one. A block-diagonal causal mask, or the per-document variable-length interfaces of fused kernels, keeps documents separate. Whether cross-document attention matters depends on data, but it should be a decision, not an accident.
- Sliding window. A causal mask further restricted to the last W keys. It caps the cache at W entries per layer, at the price of losing direct access to anything older.
Failure modes and how to catch them
- Future leakage. A missing or misaligned mask in training. Symptom: training loss far below what sampling quality suggests. Catch it with a perturbation test: change token t and assert that outputs before t do not move.
- Fully masked rows. A row whose keys are all masked, for example a padding query, gives a softmax over minus infinity everywhere and produces NaN. Using the dtype's minimum value instead avoids NaN but yields uniform attention. Either way, exclude those rows from the loss.
- Padding on the wrong side. With left padding in batched generation, position ids must skip the pad tokens or positional encodings shift. With right padding, the next token is generated after the pads unless lengths are tracked.
- Convention mix-ups. True meaning keep in one API and drop in another, or additive float masks passed where booleans were expected. Assert the expected number of visible keys per row in a unit test.
- Cache drift. Logits from cached decoding should match a full recompute of the same sequence to within numeric tolerance. Run that comparison in CI; it catches alignment, position-id and cache-indexing bugs together.
@torch.no_grad()
def assert_no_future_leak(model, vocab_size, seq_len=32, trials=5):
"""Changing token t must not change any output at positions before t."""
model.eval()
for _ in range(trials):
x = torch.randint(0, vocab_size, (1, seq_len))
t = int(torch.randint(1, seq_len, (1,)))
y = x.clone()
y[0, t] = (y[0, t] + 1) % vocab_size
a, b = model(x), model(y) # logits: (1, seq_len, vocab)
assert torch.allclose(a[:, :t], b[:, :t], atol=1e-4), f"leak at position {t}"
Choosing for a new system
If the product generates open-ended text, use a causal decoder: tooling, serving stacks and pretrained checkpoints are all built around it. If the input is complete and the output is a label, a score or a vector, start with a bidirectional encoder of modest size, which is usually cheaper per request and strong for its size; compare it with a decoder-based embedding model on your own evaluation set. If a long, fixed input conditions a generated output, consider a prefix mask or an encoder-decoder, and measure whether the bidirectional context actually improves quality enough to justify the extra complexity in caching.
What to do next
- Print the mask your model actually uses for a 6-token input and check it against the diagrams above.
- Add the future-leak perturbation test to CI for every causal model you train.
- Compare cached decoding with full recompute logits for a few prompts, and fail the build on drift.
- Audit every non-square attention call (decoding, chunked prefill, speculative verification) for bottom-right alignment.
- If you pack sequences, decide whether documents may attend to each other and enforce it with a block mask.
- Read the related pages: the KV cache, cross-attention, FlashAttention, attention sinks and the original Transformer paper.