Multi-Head Attention: Why One Spotlight Is Never Enough

Learn how Multi-Head Attention gives a Transformer multiple pairs of eyes at once—from the Detective Team analogy and tensor head-splitting to Grouped-Query Attention (GQA) in Llama-3.

26 minBeginnerCode Examples

The Core Idea: In the last lesson, we saw how a single Self-Attention spotlight lets a word look around a sentence to gather context. But what if one word needs to track three different things at the exact same time—like WHO did an action, WHERE it happened, and WHAT the tone is? If you average all three into a single spotlight, the beam gets blurry! Multi-Head Attention (MHA) fixes this by splitting the vector into HH smaller independent spotlights (Heads) that scan the sentence in parallel, each hunting for a different type of clue.

  1. Beginner: The "Blurred Spotlight" Problem (Why One Head Fails)

Look at the word "it" in this rich sentence:

"The chef dropped the glass bowl on the tile floor because it was slippery."

To truly understand the word "it" in that sentence, your brain actually asks three totally different questions at once:

Detective #1 · Noun Link
"What noun is 'it'?"

Shines a spotlight back on "glass bowl" (the object that slipped out of the chef's hands).

Detective #2 · Property Link
"How is 'it' described?"

Shines a spotlight forward on "slippery" to know why the bowl fell.

Detective #3 · Action Link
"What happened to 'it'?"

Shines a spotlight on "chef dropped" and "tile floor" to track the physical event.

Why a Single Softmax Row Cannot Do All Three Jobs Well

Remember that Softmax forces all attention percentages in a row to add up to 100%100\%—and winners take almost everything! If you only have 1 Attention Head, it might spend 90%90\% of its budget looking at "bowl" and miss "slippery" and "dropped"! By giving the layer H=8H = 8 to 6464 separate Attention Heads, each head gets its own independent 100%100\% Softmax budget to focus on a different relationship!

  1. Beginner: The Zero-Extra-Cost Trick (Splitting Sliders, Not Copying!)

Wait a minute! If we use 88 Attention Heads instead of 11, does that make the Transformer 8×8\times slower and 8×8\times heavier?

No! It costs the exact same amount of math as 1 giant head! Here is the clever trick the Transformer authors used: instead of giving all 88 heads the full 512512-number vector, they slice the 512512 sliders into 88 equal mini-chunks of 6464 sliders each!

dhead=dk=dmodelH=512 total sliders8 heads=64 sliders per head!d_{\text{head}} = d_k = \dfrac{d_{\text{model}}}{H} = \dfrac{512 \text{ total sliders}}{8 \text{ heads}} = 64 \text{ sliders per head!}
Step 1 · Split Into H Heads
512 → 8 × 64

Head #1 gets sliders 0–630\text{--}63, Head #2 gets sliders 64–12764\text{--}127, all the way to Head #8!

Step 2 · Run 8 Spotlights
8 Separate Attention Grids

Each head runs softmax(QiKiT64)Vi\text{softmax}\left(\frac{Q_i K_i^T}{\sqrt{64}}\right)V_i in parallel on its own 64-slider slice.

Step 3 · Glue Back Together
8 × 64 → 512

Concatenate all 88 mini-outputs side-by-side back into a single 512512-slider vector!

⚡ Knowledge Check

Suppose a Transformer has a model dimension of dmodel=768d_{\text{model}} = 768 (like BERT-Base and GPT-2) and uses H=12H = 12 Attention Heads. How many numbers (sliders) does each individual head work with (dkd_k)?

A) 64 sliders per head (because 768 / 12 = 64)▼
✓ Correct!Dividing 768768 equally across 1212 heads gives dk=64d_k = 64 features per head. When all 1212 heads finish, concatenating 12×6412 \times 64 restores the exact 768768-dimensional vector!
B) 9,216 sliders per head (768 * 12)▼
✕ Incorrect.Multi-Head Attention splits dmodeld_{\text{model}} across the HH heads (dk=dmodel/Hd_k = d_{\text{model}} / H) so computational cost stays constant.

  1. Medium: The Final Mixing Desk (WO\mathbf{W}_O) & Exact Equations

After the HH heads finish scanning the sentence and we glue their 6464-slider results side-by-side back into a 512512-slider vector, there is one final step that beginners often overlook!

Right now, sliders 0–630\text{--}63 only know what Detective #1 found, and sliders 64–12764\text{--}127 only know what Detective #2 found—the detectives haven't actually shared their notes with each other yet!

To let all HH detectives combine their findings, we multiply the glued vector by one final learnable weight matrix called the Output Projection Matrix (WO∈Rdmodel×dmodel\mathbf{W}_O \in \mathbb{R}^{d_{\text{model}} \times d_{\text{model}}}):

The Complete Multi-Head Attention Equation
MultiHead(Q,K,V)=Concat(head1,  head2,  …,  headH)WO\text{MultiHead}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \text{Concat}\left(\text{head}_1, \; \text{head}_2, \; \dots, \; \text{head}_H\right)\mathbf{W}_O
whereheadi=Attention(XWQ(i),  XWK(i),  XWV(i))\text{where} \quad \text{head}_i = \text{Attention}\left(\mathbf{X}\mathbf{W}_Q^{(i)}, \; \mathbf{X}\mathbf{W}_K^{(i)}, \; \mathbf{X}\mathbf{W}_V^{(i)}\right)

How PyTorch Computes All H Heads in ONE Matrix Multiply (Tensor Reshape Trick!)

In Python code, we NEVER write a slow for loop over the 88 heads! Instead, we use one giant nn.Linear(512, 512) layer and reshape the tensor in 3 steps:

1. Linear Projection
(B, T, 512)
All 8 heads packed together
2. view + transpose
(B, 8, T, 64)
Splits 512 into 8 heads of 64!
3. Parallel Attention
(B, 8, T, T)
GPU runs all 8 grids at once!

  1. Medium: What Real Attention Heads Actually Learn in the Wild

When researchers open up trained Transformers (like BERT, GPT-2, and Llama-3) and inspect what each head looks at, they find that heads automatically specialize into distinct jobs without ever being told to:

1. Previous-Word / Local Heads

Always look 11 or 22 seats to the left—acting like a local Bigram/Trigram detector to bind adjectives to nouns (like "red" → "apple").

2. Pronoun / Coreference Heads

Connect pronouns ("he", "she", "it", "they") directly back to the person or object's name across long paragraphs!

3. Induction / Copy Heads (In-Context Learning)

Search earlier in your prompt for patterns like [A][B] ... [A] and predict [B] next—powering few-shot prompting and code completion!

4. Punctuation / "Resting" Heads

Park their spotlight on the very first token (<BOS>) or a period (.) when the specific clue they hunt for isn't present in the sentence!

⚡ Knowledge Check

What is the role of the final Output Projection matrix WO\mathbf{W}_O at the very end of a Multi-Head Attention layer?

A) It mixes and combines the information discovered by all H separate heads across the full vector▼
✓ Correct!Concatenation simply places the HH head outputs side-by-side; multiplying by WO\mathbf{W}_O allows features from Head #1, Head #2, ..., Head #H to interact and blend together!
B) It converts the vector directly into English text strings▼
✕ Incorrect.WO\mathbf{W}_O outputs a dense vector of size dmodeld_{\text{model}} that stays inside the Transformer block to pass into the next layer.

  1. Advanced: Grouped-Query Attention (GQA) — How Llama-3 Saves 8x GPU Memory!

When ChatGPT or Llama-3 generates a long answer one word at a time, it saves the past Key (K\mathbf{K}) and Value (V\mathbf{V}) vectors of every earlier word in GPU memory (called the KV Cache) so it doesn't have to recompute them.

In standard Multi-Head Attention (MHA), if you have 3232 Query heads, you also have to store 3232 separate Key heads and 3232 Value heads in GPU memory—eating up tens of gigabytes of VRAM on long chats!

To solve this memory bottleneck, modern LLMs use Grouped-Query Attention (GQA) (used in Llama-3, Mistral, Gemma-2, and Qwen):

Attention VariantQuery Heads (HQH_Q)Key / Value Heads (HKVH_{KV})KV Cache Memory & Speed
1. Multi-Head Attention (MHA)32 Query Heads32 Key/Value Heads (1:1)100% KV Cache Memory (Heavy VRAM on long context)
2. Multi-Query Attention (MQA)32 Query Heads1 Shared Key/Value Head (32:1)32x smaller memory, but slight accuracy drop
3. Grouped-Query Attention (GQA)32 Query Heads8 Key/Value Heads (4:1 Groups)4x–8x less memory with 99.9% of MHA quality!
⚡ Knowledge Check

In Llama-3 8B, each layer has 3232 Query Heads (HQ=32H_Q = 32) but only 88 Key/Value Heads (HKV=8H_{KV} = 8) using Grouped-Query Attention (GQA). How does that work, and why do we do it?

A) Every group of 4 Query heads shares 1 Key/Value head, cutting KV Cache GPU memory by 4x during generation!▼
✓ Correct!Sharing 1 Key/Value head across 4 Query heads slashes the KV Cache memory footprint by 75%75\% (4×4\times smaller) while keeping all 3232 independent Query spotlights active!
B) It deletes 24 Query heads at test time▼
✕ Incorrect.All 3232 Query heads remain active; they are simply grouped into 88 teams of 44 that share the 88 Key/Value heads.

  1. Visual Explanation: The 3-Head Detective Split & Mixing Desk

Instead of a single blurry beam, watch how Multi-Head Attention slices an incoming word vector into 3 Specialist Detective Heads, lets each head lock onto a different clue in the sentence, and blends their discoveries at the WO\mathbf{W}_O Mixing Desk:

  1. INCOMING WORD VECTOR FOR "it" (d_model = 12 Sliders)
Sliced into 3 Heads of 4 Sliders Each!

Sliders [0..3] → Head 1

Sliders [4..7] → Head 2

Sliders [8..11] → Head 3

🕵️ HEAD #1 · NOUN HUNTER

d_k = 4

Asks: "Which physical object does 'it' refer to?"

➔ "glass bowl"88%
"chef"7%
"slippery"5%

🕵️ HEAD #2 · TRAIT HUNTER

d_k = 4

Asks: "What adjective describes 'it'?"

➔ "slippery"84%
"glass bowl"10%
"dropped"6%

🕵️ HEAD #3 · ACTION HUNTER

d_k = 4

Asks: "What action happened to 'it'?"

➔ "dropped on floor"81%
"chef"14%
"glass bowl"5%

🎛️ STEP 3 · CONCATENATE + W_O MIXING DESK

Output Shape Restored: (d_model = 12)

Glues [Bowl] + [Slippery] + [Dropped] back into 1212 sliders and multiplies by WO\mathbf{W}_O—so the final vector for "it" knows it is a "slippery glass bowl that got dropped on the floor"!

⚡ Click to Compare: Standard MHA vs. Llama-3 Grouped-Query Attention (GQA) Wiring

Compare Wiring

▼

Standard MHA (GPT-2 / BERT):

• 32 Query Heads
• 32 Key Heads + 32 Value Heads
Requires storing 64 full K/V tensors in GPU cache!

Grouped-Query Attention (Llama-3):

• 32 Query Heads (8 groups of 4)
• Only 8 Key Heads + 8 Value Heads!
Uses 4x less GPU RAM with zero loss in smarts!

  1. Python Implementation: Multi-Head Attention From Scratch in PyTorch

Here is a complete, runnable PyTorch implementation of Multi-Head Attention showing how professional LLM code splits d_model into num_heads using .view() and .transpose()—running all heads in parallel without a single Python loop:

multi_head_attention.pyPython 3.11+ · PyTorch 2.x
import mathimport torchimport torch.nn as nnimport torch.nn.functional as F class MultiHeadAttention(nn.Module):    def init(self, dModel: int, numHeads: int):        super().init()        assert dModel % numHeads == 0, "dModel must be divisible by numHeads!"        self.dModel = dModel        self.numHeads = numHeads        self.dHead = dModel // numHeads                  # e.g., 512 // 8 = 64 sliders per head         self.wQ = nn.Linear(dModel, dModel, bias=False)        self.wK = nn.Linear(dModel, dModel, bias=False)        self.wV = nn.Linear(dModel, dModel, bias=False)        self.wO = nn.Linear(dModel, dModel, bias=False)  # Final Mixing Desk (W_O)     def forward(self, x: torch.Tensor) -> torch.Tensor:        B, T, C = x.shape         # 1. Project Q, K, V and split C into (numHeads, dHead) -> Transpose to (B, H, T, dHead)        Q = self.wQ(x).view(B, T, self.numHeads, self.dHead).transpose(1, 2)        K = self.wK(x).view(B, T, self.numHeads, self.dHead).transpose(1, 2)        V = self.wV(x).view(B, T, self.numHeads, self.dHead).transpose(1, 2)         # 2. Run Scaled Dot-Product Attention across all H heads in parallel!        scores = (Q @ K.transpose(-2, -1)) / math.sqrt(self.dHead)      # Shape: (B, H, T, T)        weights = F.softmax(scores, dim=-1)        headOutputs = weights @ V                                       # Shape: (B, H, T, dHead)         # 3. Concatenate all H heads back into (B, T, dModel) and mix with W_O!        glued = headOutputs.transpose(1, 2).contiguous().view(B, T, C)        return self.wO(glued)                                           # Shape: (B, T, dModel) torch.manual_seed(42)mha = MultiHeadAttention(dModel=512, numHeads=8)sampleInput = torch.randn(2, 10, 512)                                   # 2 sentences, 10 words, 512 slidersoutput = mha(sampleInput)print("Input Shape:", sampleInput.shape, "-> MHA Output Shape:", output.shape)

Pro Tip (Why Do We Need .contiguous() Before .view(B, T, C)?):

Look at Step 3 in the code above: headOutputs.transpose(1, 2).contiguous().view(B, T, C). In PyTorch, calling .transpose(1, 2) swaps the axes without rearranging the numbers in RAM. Calling .contiguous() neatly lines the numbers up in memory row-by-row so .view(B, T, C) can glue the 88 heads of 6464 back into 512512 without throwing a PyTorch stride error!

Key Points

✓Multi-Head Attention (MHA) splits each token's dmodeld_{\text{model}} vector into HH smaller subspaces of size dk=dmodel/Hd_k = d_{\text{model}} / H, allowing the model to attend to multiple relationships simultaneously.
✓Because each head operates on a slice of size dmodel/Hd_{\text{model}} / H, running HH heads in parallel costs the exact same total FLOPs and parameters as running a single full-sized head.
✓After all HH heads compute their outputs in parallel, their vectors are concatenated back to size dmodeld_{\text{model}} and multiplied by the Output Projection matrix WO\mathbf{W}_O so heads can share information.
✓Different heads automatically specialize during training into roles like tracking pronouns, binding nearby adjectives, or copying patterns (Induction Heads).
✓Modern LLMs (Llama-3, Mistral, Gemma) use Grouped-Query Attention (GQA)—sharing 1 Key/Value head across every 4 to 8 Query heads to slash KV Cache GPU memory by 4–8×4\text{--}8\times.

Common Mistakes

✕ Choosing a d_model that is not evenly divisible by num_heads.

Because every head gets dmodel/Hd_{\text{model}} / H sliders, picking dmodel=512d_{\text{model}} = 512 with H=12H = 12 crashes immediately (512/12=42.66512 / 12 = 42.66). Always pick HH that divides dmodeld_{\text{model}} cleanly (like 512/8=64512 / 8 = 64 or 768/12=64768 / 12 = 64).

✕ Scaling attention scores by sqrt(d_model) instead of sqrt(d_head).

Each individual head computes a dot product over dhead=64d_{\text{head}} = 64 numbers, not dmodel=512d_{\text{model}} = 512. Always divide scores by dhead\sqrt{d_{\text{head}}} (e.g., 64=8\sqrt{64} = 8)!

✕ Calling .view(B, T, C) directly after .transpose(1, 2) without .contiguous() (or .reshape()).

Transposing the head and sequence dimensions makes the tensor non-contiguous in memory. Always call .contiguous().view(B, T, C) or .reshape(B, T, C) when concatenating heads back together.

The Big Picture

Single-Head Attention (One Blurred Spotlight)

1 Softmax Budget for All Clues → Grammar, Pronouns & Tone Get Averaged Together

Multi-Head & Grouped-Query Attention (Specialist Detective Team)

Slice Sliders Into H Heads → Parallel Spotlights Hunt Different Clues → Mix Together via W_O!

The big takeaway is simple: Multi-Head Attention gives the AI multiple pairs of eyes at zero extra math cost. While one head links a pronoun to its noun, a second head checks verb tense, and a third head tracks the overall topic—then the WO\mathbf{W}_O mixing desk blends all their answers into one rich vector.

Remember: Multi-Head Attention is the communication engine where words talk to each other—in our very next lesson, we will wrap Multi-Head Attention with Feed-Forward Networks, Residual Highways, and LayerNorm to build a complete Transformer Block!