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 smaller independent spotlights (Heads) that scan the sentence in parallel, each hunting for a different type of clue.
- 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:
Shines a spotlight back on "glass bowl" (the object that slipped out of the chef's hands).
Shines a spotlight forward on "slippery" to know why the bowl fell.
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 —and winners take almost everything! If you only have 1 Attention Head, it might spend of its budget looking at "bowl" and miss "slippery" and "dropped"! By giving the layer to separate Attention Heads, each head gets its own independent Softmax budget to focus on a different relationship!
- Beginner: The Zero-Extra-Cost Trick (Splitting Sliders, Not Copying!)
Wait a minute! If we use Attention Heads instead of , does that make the Transformer slower and 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 heads the full -number vector, they slice the sliders into equal mini-chunks of sliders each!
Head #1 gets sliders , Head #2 gets sliders , all the way to Head #8!
Each head runs in parallel on its own 64-slider slice.
Concatenate all mini-outputs side-by-side back into a single -slider vector!
Suppose a Transformer has a model dimension of (like BERT-Base and GPT-2) and uses Attention Heads. How many numbers (sliders) does each individual head work with ()?
A) 64 sliders per head (because 768 / 12 = 64)▼
B) 9,216 sliders per head (768 * 12)▼
- Medium: The Final Mixing Desk () & Exact Equations
After the heads finish scanning the sentence and we glue their -slider results side-by-side back into a -slider vector, there is one final step that beginners often overlook!
Right now, sliders only know what Detective #1 found, and sliders only know what Detective #2 found—the detectives haven't actually shared their notes with each other yet!
To let all detectives combine their findings, we multiply the glued vector by one final learnable weight matrix called the Output Projection Matrix ():
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 heads! Instead, we use one giant nn.Linear(512, 512) layer and reshape the tensor in 3 steps:
- 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:
Always look or seats to the left—acting like a local Bigram/Trigram detector to bind adjectives to nouns (like "red" → "apple").
Connect pronouns ("he", "she", "it", "they") directly back to the person or object's name across long paragraphs!
Search earlier in your prompt for patterns like [A][B] ... [A] and predict [B] next—powering few-shot prompting and code completion!
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!
What is the role of the final Output Projection matrix 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▼
B) It converts the vector directly into English text strings▼
- 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 () and Value () 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 Query heads, you also have to store separate Key heads and 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 Variant | Query Heads () | Key / Value Heads () | KV Cache Memory & Speed |
|---|---|---|---|
| 1. Multi-Head Attention (MHA) | 32 Query Heads | 32 Key/Value Heads (1:1) | 100% KV Cache Memory (Heavy VRAM on long context) |
| 2. Multi-Query Attention (MQA) | 32 Query Heads | 1 Shared Key/Value Head (32:1) | 32x smaller memory, but slight accuracy drop |
| 3. Grouped-Query Attention (GQA) | 32 Query Heads | 8 Key/Value Heads (4:1 Groups) | 4x–8x less memory with 99.9% of MHA quality! |
In Llama-3 8B, each layer has Query Heads () but only Key/Value Heads () 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!▼
B) It deletes 24 Query heads at test time▼
- 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 Mixing Desk:
- INCOMING WORD VECTOR FOR "it" (d_model = 12 Sliders)
Sliders [0..3] → Head 1
Sliders [4..7] → Head 2
Sliders [8..11] → Head 3
🕵️ HEAD #1 · NOUN HUNTER
d_k = 4Asks: "Which physical object does 'it' refer to?"
🕵️ HEAD #2 · TRAIT HUNTER
d_k = 4Asks: "What adjective describes 'it'?"
🕵️ HEAD #3 · ACTION HUNTER
d_k = 4Asks: "What action happened to 'it'?"
🎛️ STEP 3 · CONCATENATE + W_O MIXING DESK
Output Shape Restored: (d_model = 12)Glues [Bowl] + [Slippery] + [Dropped] back into sliders and multiplies by —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▼
⚡ Click to Compare: Standard MHA vs. Llama-3 Grouped-Query Attention (GQA) Wiring
▼
• 32 Query Heads
• 32 Key Heads + 32 Value Heads
Requires storing 64 full K/V tensors in GPU cache!
• 32 Query Heads (8 groups of 4)
• Only 8 Key Heads + 8 Value Heads!
Uses 4x less GPU RAM with zero loss in smarts!
- 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:
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 heads of back into without throwing a PyTorch stride error!
Key Points
Common Mistakes
✕ Choosing a d_model that is not evenly divisible by num_heads.
Because every head gets sliders, picking with crashes immediately (). Always pick that divides cleanly (like or ).
✕ Scaling attention scores by sqrt(d_model) instead of sqrt(d_head).
Each individual head computes a dot product over numbers, not . Always divide scores by (e.g., )!
✕ 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 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!