The next answer must not be visible yet
We now know the attention table: a row says how a position mixes different sources. During training, the whole example sentence is already in memory. Even so, a position may only read itself and earlier positions. Otherwise it could simply copy the next token it is supposed to predict.
For three positions, the allowed sources are: position 1 reads only 1; position 2 reads 1 and 2; position 3 reads 1, 2 and 3. A mask is a table that records exactly these permissions. "Causal" here means: a computation does not depend on later text positions.
To exclude a source, you set its score to minus infinity, written −∞, before softmax. This is a mathematical limit notation: the exponential function approaches zero for ever more negative numbers. So the forbidden source gets weight zero. In code, a truth value like True or False can describe the permission; which direction means "allowed" depends on the specific function.
Computing in parallel without seeing the future
During training, all tokens of a sequence are available as an array. Even so, the output at position may only read input positions up to and including . Its target is ; that sits one position further to the right. The causal mask forbids exactly this look to the right.
This does not mean we have to run a separate forward pass in a loop for every prefix. We compute many query rows at once and mask the disallowed entries. This way the model gets a correctly restricted prefix at every position and can learn all the corresponding next-token targets in parallel.
Why the diagonal is allowed
The input at position is . The corresponding target is . So the current input token belongs to the allowed context. A mask that also forbids the diagonal leaves the first position with no allowed source at all: all scores would be −∞, and softmax would not give a valid distribution (NaN in PyTorch). Our standard decoder allows the diagonal.
T = q.size(-2)
allowed = torch.ones(T, T, dtype=torch.bool, device=q.device).tril()
scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
scores = scores.masked_fill(~allowed, float('-inf'))
weights = scores.softmax(dim=-1)
out = weights @ vForbidden scores are set to minus infinity. The exponential function turns them into zero. A plain score of zero would be wrong: softmax can assign it a positive probability.
A dangerous mask convention
Different PyTorch APIs use boolean masks differently. In scaled_dot_product_attention, a true boolean mask position means "may take part". In other attention interfaces, true can mean "blocked". So never copy a mask blindly between APIs.
For the complete equal-length decoder forward pass, we use:
out = F.scaled_dot_product_attention(
q, k, v,
dropout_p=0.0,
is_causal=True,
)You must set the dropout probability to zero explicitly if you don't want dropout. The functional API does not infer this automatically from model.eval(). The reference model never uses attention dropout.
Multi-head attention
D is still the total representation width. H is the number of parallel attention heads, and lowercase d is the width of a single head. With D = 128 and H = 4, d = 32. Instead of using one head of width , we use heads of width . Each head has its own query, key and value projection. In practice, we can compute all heads in one large linear projection and then reshape the head axis out of it.
q = self.q_proj(x) # [B,T,D]
q = q.view(B, T, H, d).transpose(1, 2) # [B,H,T,d]
# same for k and v
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
out = out.transpose(1, 2).contiguous().view(B, T, D)
out = self.out_proj(out)The heads mix information separately, but are then merged. An output projection can combine their results. At the same total width, "more heads" does not simply lead to a proportionally larger parameter count in the standard projections. Instead, the individual heads become narrower.
The prefix test
Create two sequences that are identical up to position 4 and contain different tokens after that. In a causal model, the logits up to position 4 must stay the same, apart from numerical rounding noise. Put the model in eval mode and switch off sources of randomness such as dropout.
model.eval()
a = torch.tensor([[1, 5, 6, 7, 8, 9]])
b = torch.tensor([[1, 5, 6, 7, 8, 3]])
with torch.no_grad():
la, lb = model(a), model(b)
assert torch.allclose(la[:, :5], lb[:, :5], atol=1e-5)Our test project uses this idea. This test tells you far more than glancing at a loss. A missing mask can improve the training loss a lot, precisely because it makes the task unfairly easy.
Causality also applies to other layers
In our model, LayerNorm or RMSNorm only normalize the last feature axis of a position. If you accidentally normalized over the time axis, future positions could leak into the statistics. Data preprocessing that writes labels into the input can also bypass causal attention. A correct attention block is therefore not proof for the whole model; test the complete model.
The browser experiment
Click a query row in the matrix. It shows the allowed positions, not learned relationships. Switch off the mask and look at what extra information would then give away the training target. The matrix is deliberately labeled as a structural view: the model does not have to weight an allowed connection strongly.