Scaled Dot-Product Attention
Goal
Explain why attention scales query-key dot products before softmax, give an intuition for why the standard factor is sqrt(d_k), and test how the scale changes probability sharpness.
The problem: wider heads can make softmax too sharp
L5.4 showed that a dot product adds one product per vector component. If a head has more components, its raw query-key scores can naturally spread farther from zero even before that wider head has learned a more confident relationship.
Softmax reacts strongly to score gaps. If the logits become unnecessarily far apart, the largest one can receive almost all the probability. We want to control that width-dependent scale without changing which key scored highest.
That is why attention divides the query-key scores before softmax. When all query vectors are stacked into Q and all key vectors into K, QK^T computes the query-key dot products together. The full operation is:
softmax(QK^T / sqrt(d_k)) V
Read it in order:
query-key scores
→ divide by sqrt(d_k)
→ softmax into weights
→ mix values
The division turns the score “volume” down; it does not reorder the scores.
Hold the logits fixed and isolate the scaling effect
This controlled comparison uses the same synthetic raw logits as the Lab. Change only the denominator dimension and watch probability sharpness change without changing the score order.
| Position | Fixed raw logit | Scaled logit | Softmax probability |
|---|---|---|---|
| 1 | 0.000 | 0.000 | 9.0% |
| 2 | 4.000 | 2.000 | 66.5% |
| 3 | 2.000 | 1.000 | 24.5% |
Predict
Why the square root appears
Here is the useful intuition behind sqrt(d_k).
A query-key dot product is a sum of d_k component products:
q1*k1 + q2*k2 + ... + q_d*k_d
Near a typical random initialization, imagine those component products as small positive and negative contributions centered roughly around zero. They do not all push in the same direction, so the usual size of the sum does not grow like d_k itself.
For roughly independent, similarly scaled contributions:
- the variance of the sum grows roughly in proportion to
d_k; - the typical distance from zero—the standard deviation—is the square root of variance;
- so the typical score magnitude grows roughly like
sqrt(d_k).
Dividing by sqrt(d_k) counters that growth and keeps the score scale more comparable as the head width changes.
You do not need to derive probability theory here. The important chain is:
more components
→ wider raw dot-product spread
→ softer/harder softmax behavior can change just because width changed
→ divide by the typical sqrt(d_k) growth
→ keep scores in a more stable range
This is a statistical initialization intuition, not a promise that every real query-key vector has exactly that variance.
Compare raw and scaled scores in the Lab
The Lab deliberately holds the raw score pattern fixed at:
raw = [0.0, 4.0, 2.0]
and applies the denominator for head_dim values 1, 4, and 16.
- Click Run unchanged.
- Compare the lines for
head_dim 1,4, and16. With the same raw scores, the scaled gaps become smaller assqrt(head_dim)grows, so the softmax distribution becomes less concentrated. - Find
raw = [0.0, 4.0, 2.0]and change it to:
raw = [0.0, 8.0, 4.0]
- Before running, predict that doubling all score gaps will make each corresponding softmax distribution sharper.
- Click Run and compare probabilities for the same
head_dimbefore and after the edit. - Restore
raw = [0.0, 4.0, 2.0].
Loading lab…
The Lab isolates the effect of the denominator. Because it supplies a fixed synthetic score list rather than constructing random high-dimensional query/key vectors, it does not by itself demonstrate the statistical sqrt(d_k) growth argument above. The explanation and the experiment answer two related but different questions.
Scaling does not reorder the keys
All scores in one row are divided by the same positive number. That preserves their ordering.
If key A scored higher than key B before scaling, it still scores higher after scaling. What changes is how far apart the scores appear to softmax and therefore how concentrated the normalized weights become.
Quick Check
Explain it back
Explain the difference between these two claims:
- “Scaling makes every attention relationship weaker.” — not correct; ranking is preserved.
- “Scaling keeps score magnitude from growing just because head width grows, which keeps softmax in a more useful numerical regime.” — the intended idea.
Key Takeaways
- Query-key dot products are scaled before softmax.
- Under a common initialization intuition, raw dot-product variance grows with
d_k, so typical magnitude grows likesqrt(d_k). - Dividing by
sqrt(d_k)counters that width-dependent growth. - Positive scaling preserves score ordering but changes distribution sharpness.
- Separate the statistical reason for the factor from the Lab that demonstrates its numerical effect.
Next Lesson
Next, convert the normalized scores into a weighted mixture of value vectors.
References
- Vaswani et al., Attention Is All You Need.
- PyTorch, scaled_dot_product_attention.
- Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.
Completion is stored locally on this device.