Causal Attention: The Autoregressive Mask
Introduction
In self-attention, every token can look at every other token. But in autoregressive models (language models generating text one token at a time), a token should never see future tokens—it violates the causality assumption that the model is predicting the next word given only past context.
Causal attention enforces this by applying a triangular mask: set all scores for future positions to $-\infty$ before softmax. This post implements the mask and shows how torch::tril, torch::ones, and masked_fill_ work together in LibTorch.
The Causal Principle
In a sequence $[\text{tok}_0, \text{tok}_1, \text{tok}_2, \text{tok}_3]$:
In matrix form, the attention score matrix $S$ is $(T \times T)$. We want to zero out the upper triangle (future positions). After softmax, those positions contribute zero weight.
LibTorch Implementation
Mask Construction Step-by-Step
torch::ones({T, T}, x.options()) creates a tensor of all 1s with shape $(T \times T)$, on the same device (CPU/GPU) and dtype as the input.
torch::tril(...) keeps only the lower triangle (including diagonal) and zeros the upper triangle. For a $4 \times 4$ matrix, it produces:
masked_fill_ replaces positions where the condition is true with a fill value. Here, causal_mask == 0 is true in the upper triangle, so we fill those with $-\infty$ (represented by std::numeric_limits<float>::lowest()).
| After Step | Shape | Content |
|---|---|---|
| torch::ones | (T, T) | All 1s |
| torch::tril | (T, T) | Lower triangle = 1, upper triangle = 0 |
| masked_fill_ | (T, T) | Lower triangle = score, upper triangle = $-\infty$ |
| torch::softmax | (T, T) | Lower triangle = softmax weight, upper triangle ≈ 0 |
Interactive Mask Visualization
▶ Building the Causal Mask
Watch how the mask is constructed step by step.
Complete Shape Trace
| Step | Variable | Shape | Notes |
|---|---|---|---|
| Input | x | (batch, T, d_model) | Token embeddings |
| Project Q, K, V | Q, K, V | (batch, T, d_k) | Linear projections |
| Raw scores | scores | (batch, T, T) | Q @ K^T / sqrt(d_k) |
| Create ones | temp | (T, T) | torch::ones |
| Apply tril | causal_mask | (T, T) | Lower triangle = 1, upper = 0 |
| Masked fill | scores | (batch, T, T) | Upper triangle set to -∞ |
| Softmax | attn_weights | (batch, T, T) | Rows normalized, future columns ≈ 0 |
| Output | output | (batch, T, d_k) | Weighted sum of values |
Key Insights
masked_fill_, broadcasting automatically extends the mask across the batch.
What Comes Next
Causal attention applies a static mask to a single source (self-attention). But in encoder-decoder models (translation, summarization), we often need cross-attention: the decoder attends to the encoder output, not itself. Next: Cross-Attention.