attention head scaling

**Attention Head Scaling** is the **sqrt(d_k) divisor used inside scaled dot-product attention so scores remain in a numerically stable range before the softmax** — dividing dot products by the square root of the key dimension prevents very large values that would collapse softmax probabilities and choke gradients. **What Is Head Scaling?** - **Definition**: The factor 1/sqrt(d_k) applied to the QK^T result before the softmax step in multi-head attention. - **Key Feature 1**: Without scaling, dot products grow with d_k, making softmax saturate and gradients vanish. - **Key Feature 2**: Scaling keeps logits around zero, so the softmax spreads attention weight across tokens. - **Key Feature 3**: The same scalar is applied to every head, keeping relative relationships comparable across heads. - **Key Feature 4**: Some proposals extend scaling to additive biases or head-dependent factors. **Why Scaling Matters** - **Stability**: Prevents overflow in softmax when d_k is large. - **Gradient Flow**: Maintains non-zero gradients by avoiding saturated attention scores. - **Uniform Behavior**: Keeps the attention distribution consistent across architecture variations that change d_k. - **Theoretical Basis**: Derived from variance considerations: dot product variance equals d_k, so scaling rescales to unit variance. - **Hyperparameter Simplicity**: Makes the behavior of attention predictable across head counts and dimensions. **Scaling Variants** **Standard sqrt(d_k)**: - Default in classic Transformer models. - Works across language and vision tasks. **Head-wise Scaling**: - Each head learns its own scale via a parameter. - Helps if heads have different dimensionalities or roles. **Bias + Scale**: - Adds learnable biases to center the logits after scaling. - Useful when attention logits need calibration. **How It Works / Technical Details** **Step 1**: After computing the dot product between queries and keys, multiply the result by the scalar 1/sqrt(d_k) to normalize variance. **Step 2**: Feed the scaled logits into softmax, ensuring the distribution stays smooth and gradient-friendly; head-wise scaling further trains these scalars. **Comparison / Alternatives** | Aspect | Scaled Attention | Unscaled | Learnable Scale | |--------|------------------|----------|-----------------| | Variance Control | Yes | No | Yes | Gradient Stability | High | Low | High | Complexity | Minimal | Minimal | Slightly higher | ViT Best Practice | Required | Not recommended | Optional **Tools & Platforms** - **PyTorch / TensorFlow**: Scaling built into their multi-head attention APIs. - **timm**: Allows overriding the scaling factor for experiments. - **Custom Modules**: Implement fixed or learnable scaling by multiplying the logits tensor. - **Profiling**: Check gradient norms with vs without scaling to highlight its importance. Attention head scaling is **the simple divisor that makes multi-head attention numerically tame despite large key dimensions** — without it, the softmax becomes brittle and transformers lose their ability to learn.

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account