All concepts

Multi-Head Attention

Run several attention patterns in parallel so tokens can attend for different reasons.

Transformers & LLMs · Intermediate · ~8 min

In plain English

Run several attention passes side by side, each free to specialize — one tracks grammar, another tracks who's who — then stitch their answers together.

Why it's worth your time

It's why one attention layer can follow syntax, coreference and topic at the same time.

If you remember three things

  • Heads split the dimension, they don't multiply the compute
  • Different heads learn genuinely different relationships
  • Outputs are concatenated and mixed by one more projection

Overview

Runs several attention operations in parallel, each with its own learned Q/K/V projections, so a token can gather different kinds of context at once. The heads' outputs are concatenated and linearly projected back to the model dimension: MHA(X)=Concat(head_1...head_h) W_o.

How it works

  1. Start: Token Vectors The same input representation feeds multiple learned projections.
  2. Token Vectors -> Attention Heads Each head has its own Q, K, V projections and can focus on syntax, coreference, position, or entities.
  3. Attention Heads -> Different Patterns One head may track subject-verb links while another tracks delimiters or recency.
  4. Different Patterns -> Concat Head outputs are concatenated and projected back to model dimension.
  5. Concat -> Context Vectors The token now carries several kinds of contextual evidence.

In an interview

Instead of one attention function, multi-head attention splits the model dimension across h heads, each computing scaled dot-product attention with independent Q, K, V projections. Different heads learn different relationships — syntax, coreference, positional patterns — then their outputs are concatenated and projected by W_o. This lets one layer attend for several reasons simultaneously.

Production defaults

Head count
head_dim of 64–128 is the convention; heads = model_dim / head_dim
Pruning
many heads are redundant and can be removed with little loss — worth checking at serving time
Don't
add heads without shrinking head_dim; you'll change the parameter count, not the capability

What breaks

  • More heads didn't help — You reduced per-head dimension below what the relationship needs. Head_dim under ~32 usually degrades.
  • KV cache memory dominates at serving — Every head stores its own K and V. That's what grouped-query attention exists to fix.

Watch it explained

What are Transformers (Machine Learning Model)? — IBM Technology, 5:50

Related