← Back to Table of Contents

Chapter 4 β€” Attention Deep Dive: SDPA & Multi-Head Attention

β€œThe key innovation of attention is allowing the model to dynamically focus on relevant parts of the input, rather than compressing everything into a fixed-size vector.”

Scaled Dot-Product Attention (SDPA)

Attention is the mechanism that lets each token look at every other token in the sequence and decide what’s relevant. At its core, it’s a soft lookup: queries ask questions, keys advertise content, and values hold the actual information.

Given three matrices β€” Query (Q), Key (K), and Value (V) β€” scaled dot-product attention computes:

\[\text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^T}{\sqrt{d_k}}\right) V\]
SDPA β€” Step by Step
1. Compute scores: QK^T [B, H, T, T]
2. Scale by 1/√d_k prevents softmax saturation
3. Apply causal mask βˆ’βˆž for future positions
4. Softmax β†’ attention weights [B, H, T, T], rows sum to 1
5. Weighted sum of values weights Γ— V β†’ [B, H, T, d_k]

Tensor Shapes Through SDPA

SDPA Tensor Shapes
Q (queries)
[ B , H , T , d_k ]
K (keys)
[ B , H , T , d_k ]
QK^T (scores)
[ B , H , T , T ]
← This is O(TΒ²) in memory!
V (values)
[ B , H , T , d_k ]
Output
[ B , H , T , d_k ]

Why Scale by √d_k?

Without scaling, the dot product of Q and K grows proportionally to d_k. Large dot products push softmax into saturated regions where gradients vanish. Dividing by √d_k keeps the variance of the scores at ~1 regardless of dimension:

If $q_i, k_i \sim \mathcal{N}(0, 1)$, then $\text{Var}(q \cdot k) = d_k$. After scaling: $\text{Var}!\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = 1$.

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
import torch
import torch.nn.functional as F

def scaled_dot_product_attention(Q, K, V, mask=None):
    """
    Q, K, V: [B, H, T, d_k]
    mask: [T, T] or [B, 1, T, T] β€” True means IGNORE
    Returns: [B, H, T, d_k]
    """
    d_k = Q.shape[-1]
    scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)  # [B, H, T, T]

    if mask is not None:
        scores = scores.masked_fill(mask, float('-inf'))

    weights = F.softmax(scores, dim=-1)  # [B, H, T, T]
    output = torch.matmul(weights, V)    # [B, H, T, d_k]
    return output, weights

Causal Masking

In decoder-only models (GPT, LLaMA), each token can only attend to itself and previous tokens β€” it must not see the future. This is enforced with a causal mask: an upper-triangular matrix of βˆ’βˆž values that, after softmax, zero out attention to future positions.

Causal (Lower-Triangular) Attention Mask
Mask Matrix (T=5)
  • βœ“ Β· Β· Β· Β·
  • βœ“ βœ“ Β· Β· Β·
  • βœ“ βœ“ βœ“ Β· Β·
  • βœ“ βœ“ βœ“ βœ“ Β·
  • βœ“ βœ“ βœ“ βœ“ βœ“
After Softmax
  • 1.0 0 0 0 0
  • 0.6 0.4 0 0 0
  • 0.2 0.3 0.5 0 0
  • 0.1 0.2 0.3 0.4 0
  • 0.1 0.1 0.2 0.3 0.3
1
2
3
4
5
6
7
8
9
10
def causal_mask(T, device='cpu'):
    """Returns a boolean mask: True = masked (ignored)."""
    return torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1)

mask = causal_mask(5)
# tensor([[False,  True,  True,  True,  True],
#         [False, False,  True,  True,  True],
#         [False, False, False,  True,  True],
#         [False, False, False, False,  True],
#         [False, False, False, False, False]])

Multi-Head Attention (MHA)

Scaled dot-product attention showing Q/K/V projections, causal attention weight matrix, and multi-head concatenation
Scaled dot-product attention with causal masking and multi-head mechanism

Instead of computing a single attention function, multi-head attention runs H parallel attention operations (called β€œheads”), each on a different d_k-dimensional subspace. This lets different heads learn different types of relationships (syntactic, semantic, positional, etc.).

Multi-Head Attention β€” Split, Attend, Concatenate
Input x: [B, T, d_model]
Linear projections: W_Q, W_K, W_V each [d_model, d_model] β†’ Q, K, V: [B, T, d_model]
Reshape to [B, H, T, d_k] d_k = d_model / H
H parallel SDPA heads each: [B, 1, T, d_k] β†’ [B, 1, T, d_k]
Concatenate β†’ [B, T, d_model] reshape back
Output projection W_O [d_model, d_model] β†’ [B, T, d_model]

For LLaMA-3-8B: d_model = 4096, H = 32 heads, d_k = 4096/32 = 128 per head.

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, n_heads):
        super().__init__()
        self.n_heads = n_heads
        self.d_k = d_model // n_heads

        self.W_q = nn.Linear(d_model, d_model, bias=False)
        self.W_k = nn.Linear(d_model, d_model, bias=False)
        self.W_v = nn.Linear(d_model, d_model, bias=False)
        self.W_o = nn.Linear(d_model, d_model, bias=False)

    def forward(self, x, mask=None):
        B, T, _ = x.shape

        # Project and reshape: [B, T, d_model] β†’ [B, H, T, d_k]
        Q = self.W_q(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
        K = self.W_k(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
        V = self.W_v(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)

        # Scaled dot-product attention per head
        attn_out, _ = scaled_dot_product_attention(Q, K, V, mask=mask)  # [B, H, T, d_k]

        # Concatenate heads: [B, H, T, d_k] β†’ [B, T, d_model]
        concat = attn_out.transpose(1, 2).contiguous().view(B, T, -1)

        # Final projection
        return self.W_o(concat)  # [B, T, d_model]

Self-Attention vs Cross-Attention

Self-Attention vs Cross-Attention
Self-Attention
  • Q, K, V all come from the same sequence
  • Each token attends to all other tokens in the sequence
  • Used in both encoder and decoder
  • Decoder uses causal masking
Cross-Attention
  • Q from decoder, K and V from encoder
  • Decoder tokens attend to encoder representations
  • Used in encoder-decoder models (T5, BART)
  • Not present in decoder-only models (GPT, LLaMA)

In self-attention, the query, key, and value all come from the same input. In cross-attention, Q comes from one sequence (decoder) while K and V come from another (encoder output). Decoder-only models like GPT and LLaMA use only self-attention.

Flash Attention

The attention score matrix [B, H, T, T] is O(TΒ²) in memory. For T = 128K tokens with 32 heads and FP16, that’s 128K Γ— 128K Γ— 32 Γ— 2 bytes = 1 TB β€” far too large to materialize. Flash Attention (Dao et al., 2022) solves this by:

  1. Tiling: Compute attention in small blocks that fit in GPU SRAM (shared memory)
  2. Online softmax: Compute softmax incrementally without materializing the full score matrix
  3. Kernel fusion: Fuse the matmul, softmax, and output matmul into a single GPU kernel
Flash Attention β€” Tiled Computation
Standard Attention
Materializes full TΓ—T matrix in HBM. O(TΒ²) memory. Multiple kernel launches.
Flash Attention
Processes TΓ—T in tiles within SRAM. O(T) memory. Single fused kernel. 2–4Γ— faster.

In practice, you almost never implement attention yourself. PyTorch provides an optimized implementation that auto-selects the best backend (Flash Attention, memory-efficient attention, or math fallback):

1
2
3
4
5
6
7
8
9
import torch.nn.functional as F

# PyTorch's optimized SDPA β€” automatically uses Flash Attention when available
output = F.scaled_dot_product_attention(
    Q, K, V,                    # [B, H, T, d_k]
    attn_mask=None,
    is_causal=True,             # applies causal mask internally
    dropout_p=0.0,
)  # β†’ [B, H, T, d_k]

Attention Head Specialization

Different heads learn to attend to different things. Research has shown that transformer heads naturally specialize:

Head Type What It Learns Example
Positional Attend to nearby tokens β€œthe β†’ cat” (adjacent word)
Syntactic Subject-verb agreement β€œThe cats β†’ are” (long-range grammar)
Copying Attend to identical or similar tokens Repetitions, references
Induction [A][B]…[A] β†’ predict [B] In-context learning patterns
Rare/dead Low-entropy, nearly uniform Some heads are redundant

Understanding head specialization is important for attention variant design β€” some heads are more important than others, which motivates the grouped and multi-query approaches in Chapter 5.

What’s Next

Standard multi-head attention gives each head its own Q, K, and V projections β€” but this creates a memory bottleneck when caching K and V during inference. The next chapter explores how modern models address this with grouped-query attention (GQA), multi-query attention (MQA), and multi-latent attention (MLA).

← Previous: Chapter 3 β€” The Transformer Β· Next: Chapter 5 β€” GQA, MQA & MLA β†’


Last updated: April 2026