# How Causal Masking Works in Multi-Head Attention: Implementation Deep Dive

> Understand causal masking in multi-head attention. Learn how LLMs-from-scratch implements this autoregressive generation technique by blocking future tokens. Deep dive into the implementation.

- Repository: [Sebastian Raschka/LLMs-from-scratch](https://github.com/rasbt/LLMs-from-scratch)
- Tags: deep-dive
- Published: 2026-05-12

---

**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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/gpt2.py) at lines 56-58, the code creates a boolean mask where `True` indicates illegal positions:

```python
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:

```python
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:

```python
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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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:

```python
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:

```python
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:

```python
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:

```python
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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/pkg/llms_from_scratch/kv_cache/gpt2.py)**: Contains the baseline `MultiHeadAttention` class with standard triangular causal masking using `torch.triu` and `masked_fill_`.
- **[`ch04/03_kv-cache/gpt_with_kv_cache_optimized.py`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/pkg/llms_from_scratch/ch03.py)**: Provides `MultiHeadAttentionWrapper` with Flash Attention support via the `is_causal=True` parameter.

## 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)`, where `True` values 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 `MultiHeadAttention` classes 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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/pkg/llms_from_scratch/ch03.py) where the `MultiHeadAttentionWrapper` class forwards this parameter to leverage Flash Attention or memory-efficient attention kernels when available.