How Causal Masking Works in Multi-Head Attention: Implementation Deep Dive
Causal masking enforces autoregressive generation by blocking attention to future tokens through an upper-triangular boolean mask that sets prohibited scores to negative infinity before the softmax operation, as implemented in the rasbt/LLMs-from-scratch repository.
In transformer-based language models, causal masking in multi-head attention ensures tokens only attend to previous and current positions, preserving the left-to-right generation order and preventing information leakage. The rasbt/LLMs-from-scratch repository provides clean, educational PyTorch implementations demonstrating both standard training and optimized inference with KV-caching. This article examines the exact mechanics of the causal mask and traces its implementation across the codebase.
Why Causal Masking is Required
Autoregressive language models generate text token-by-token, where each new token depends only on preceding context. Without masking, the attention mechanism would allow queries to attend to future keys, effectively letting the model "cheat" by looking ahead at the target sequence during training. This violates the causal structure required for generation and breaks the temporal dependencies that make sequential prediction possible.
The Causal Mask Structure
A causal mask is a binary upper-triangular matrix of shape (L, L) where L is the sequence length. Entries where the column index exceeds the row index (j > i) are marked as illegal:
[[0, 1, 1, …, 1],
[0, 0, 1, …, 1],
[0, 0, 0, …, 1],
…
[0, 0, 0, …, 0]]
When applied to raw attention scores, positions marked with 1 (or True) are set to negative infinity (-inf). The subsequent softmax converts these to zero probability, ensuring no information flows from future tokens to the current position.
Implementation in the LLMs-from-scratch Repository
The repository implements causal masking across several files, with the core logic residing in pkg/llms_from_scratch/kv_cache/gpt2.py and optimized variants in the chapter-specific examples.
Creating the Boolean Mask with torch.triu
The standard construction uses PyTorch's torch.triu function to generate the upper-triangular matrix. In gpt2.py at lines 56-58, the code creates a boolean mask where True indicates illegal positions:
causal_mask = torch.triu(
torch.ones(seq_len, seq_len, dtype=torch.bool, device=x.device),
diagonal=1
)
Setting diagonal=1 excludes the main diagonal, allowing tokens to attend to themselves—a necessary property for self-attention—while blocking only strictly future positions.
Broadcasting to Batch and Head Dimensions
Before application, the 2D mask must align with the 4D attention score tensor of shape (batch, heads, query_len, key_len). The implementation broadcasts the mask via slicing and unsqueezing:
causal_mask = causal_mask[:, -num_tokens:][None, None, :, :]
This operation found at line 58 expands the mask to shape (1, 1, L, L), enabling automatic broadcasting across the batch and head dimensions during the masking operation.
Applying the Mask to Attention Scores
The actual masking occurs in-place using masked_fill_ at line 64:
attn_scores.masked_fill_(causal_mask, -torch.inf)
This replaces all upper-triangular values with -inf. The subsequent softmax computation (torch.softmax(attn_scores / sqrt(d_k), dim=-1)) then produces zero attention weights for these positions, effectively implementing the causal constraint.
Optimizing for KV-Cache Inference
When using a KV-cache for efficient generation, the causal mask must account for previously cached tokens. The optimized implementation in ch04/03_kv-cache/gpt_with_kv_cache_optimized.py (lines 96-104) dynamically constructs the mask rather than materializing a full (K, K) matrix:
if num_tokens == K: # no cache yet
causal_mask = torch.triu(torch.ones(num_tokens, K, device=x.device, dtype=torch.bool), diagonal=1)
else: # cache present
offset = K - num_tokens # tokens already stored
row_idx = torch.arange(num_tokens, device=x.device).unsqueeze(1)
col_idx = torch.arange(K, device=x.device).unsqueeze(0)
causal_mask = row_idx + offset < col_idx
The offset calculation (offset = K - num_tokens) shifts the causal boundary to account for cached key-value pairs, ensuring new tokens cannot attend to positions beyond the current generation step while maintaining compatibility with cached history.
Practical Code Examples
Minimal Causal Masking Demonstration
This standalone example shows the complete causal attention flow without library dependencies:
import torch
def causal_mask(seq_len, device):
# Upper-triangular mask (True = illegal future positions)
return torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool, device=device), diagonal=1)
# Dummy inputs
x = torch.randn(2, 5, 64) # (batch, seq_len, embed_dim)
queries = x @ torch.randn(64, 64)
keys = x @ torch.randn(64, 64)
values = x @ torch.randn(64, 64)
# Reshape for multi-head (single head for demo)
b, L, d = queries.shape
queries = queries.view(b, L, 1, d).transpose(1, 2) # (b, 1, L, d)
keys = keys.view(b, L, 1, d).transpose(1, 2)
values = values.view(b, L, 1, d).transpose(1, 2)
# Create and broadcast mask
mask = causal_mask(L, queries.device)[None, None, :, :] # (1, 1, L, L)
# Compute masked attention
scores = torch.matmul(queries, keys.transpose(-2, -1))
scores.masked_fill_(mask, -torch.inf)
weights = torch.softmax(scores / d**0.5, dim=-1)
output = torch.matmul(weights, values).transpose(1, 2) # (b, L, d)
Using the Repository's MultiHeadAttention Class
For production use, import the optimized implementation from the package:
from pkg.llms_from_scratch.kv_cache.gpt2 import MultiHeadAttention
import torch
batch, seq, dim = 2, 7, 128
x = torch.randn(batch, seq, dim)
att = MultiHeadAttention(
d_in=dim,
d_out=dim,
context_length=seq,
dropout=0.0,
num_heads=4,
qkv_bias=True
)
out, _ = att(x) # Shape: (batch, seq, dim), respecting causal constraints
Inference with KV-Cache Support
The class automatically handles causal mask offsets when caching is enabled:
from pkg.llms_from_scratch.kv_cache.gpt2 import MultiHeadAttention
att = MultiHeadAttention(128, 128, 1024, 0.0, num_heads=4)
# Initial forward pass (no cache)
x1 = torch.randn(1, 5, 128)
y1, cache = att(x1, use_cache=True)
# Subsequent generation step with offset masking
x2 = torch.randn(1, 3, 128)
y2, cache = att(x2, use_cache=True, start_pos=5, cache=cache)
The second call applies the offset-aware masking logic internally, ensuring tokens in x2 cannot attend to positions beyond index 5 within the cached context.
Key Implementation Files
pkg/llms_from_scratch/kv_cache/gpt2.py: Contains the baselineMultiHeadAttentionclass with standard triangular causal masking usingtorch.triuandmasked_fill_.ch04/03_kv-cache/gpt_with_kv_cache_optimized.py: Implements cache-aware masking with dynamic offset calculation for memory-efficient autoregressive generation.ch04/01_main-chapter-code/gpt.py: Demonstrates integration of causal attention into full transformer blocks without KV-caching.pkg/llms_from_scratch/ch03.py: ProvidesMultiHeadAttentionWrapperwith Flash Attention support via theis_causal=Trueparameter.
Summary
- Causal masking prevents information leakage by ensuring tokens only attend to previous and current positions, maintaining the autoregressive property essential for language generation.
- The mask is implemented as a boolean upper-triangular matrix constructed with
torch.triu(..., diagonal=1), whereTruevalues mark illegal future positions. - Mask application uses
masked_fill_(causal_mask, -torch.inf)before softmax to zero out future attention weights. - KV-cache optimization requires dynamic mask construction using an offset (
K - num_tokens) to shift the causal boundary when reusing cached key-value pairs. - The rasbt/LLMs-from-scratch repository implements these patterns in
MultiHeadAttentionclasses with explicit handling for both training and cached inference modes.
Frequently Asked Questions
What happens if you don't use causal masking in multi-head attention?
Without causal masking, the model can attend to future tokens in the sequence during training, causing information leakage where the prediction for position i directly sees the target token at position i+1. This prevents the model from learning proper autoregressive generation and results in invalid perplexity scores during evaluation.
Why does the implementation use diagonal=1 instead of diagonal=0?
Setting diagonal=1 in torch.triu excludes the main diagonal from the mask, allowing tokens to attend to themselves. This self-attention is necessary because a token's own representation contains crucial information for its own prediction; diagonal=0 would mask the current position itself, degrading model performance.
How does the KV-cache affect the causal mask calculation?
When using a KV-cache, previously computed keys and values are stored from earlier generation steps. The causal mask must shift its upper-triangular boundary by the number of cached tokens (the offset) so that new queries only mask future positions relative to the entire sequence history, not just the current chunk.
Can causal masking be implemented using PyTorch's native scaled_dot_product_attention?
Yes, PyTorch's torch.nn.functional.scaled_dot_product_attention supports causal masking natively via the is_causal=True flag, which internally applies an optimized causal mask. The repository demonstrates this in pkg/llms_from_scratch/ch03.py where the MultiHeadAttentionWrapper class forwards this parameter to leverage Flash Attention or memory-efficient attention kernels when available.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →