Softmax Attention: Numerically Stable Softmax
Introduction
Softmax is the function that turns attention scores into weights. At first glance, it seems simple: $\text{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}$. But computing this naively leads to overflow/underflow: if any input is large, $e^{x}$ explodes; if all inputs are small, $e^{x}$ underflows to zero.
This post breaks down the numerically stable version: subtract the maximum before exponentiating, sum the exponents, and divide. We'll implement it in C++, compare it to torch::softmax, and trace the three steps with an interactive animation.
Why Naive Softmax Fails
The textbook definition:
The problem: if $x = [1000, 1001, 999]$, then $e^{1000}$ is beyond float64 range. The function overflows and returns NaN. Even if the inputs fit, the ratio can be unstable due to floating-point rounding errors.
The Numerically Stable Trick
Use the log-sum-exp trick: subtract the maximum before exponentiating.
C++ LibTorch Implementation
The Three Steps Explained
| Step | Operation | LibTorch Code | Purpose |
|---|---|---|---|
| 1. Shift | $x_i - \max(x)$ | torch::max(x, -1, true).values | Subtract max to avoid overflow |
| 2. Exp & Sum | $e^{x_i - c}$ and $\sum e^{x_j - c}$ | torch::exp(x_shifted) + torch::sum(..., -1, true) | Compute exponents and their sum |
| 3. Divide | $\frac{e^{x_i - c}}{\sum e^{x_j - c}}$ | exp_x / sum_exp | Normalize to get probabilities |
In the code, the second argument to torch::max and torch::sum is the dimension (-1 for last), and the third argument is true for keepdim. This preserves the shape: if input is (batch, T, T), max returns (batch, T, 1), which broadcasts correctly when we subtract.
Stable vs. torch::softmax
LibTorch's built-in torch::softmax already uses numerical stabilization internally. For production code, prefer it. But understanding the manual implementation is valuable:
The manual version is slower (four separate operations instead of one fused kernel), but it shows the underlying mathematics. Most modern frameworks fuse these operations for performance.
Interactive Softmax Steps
▶ Three Steps of Numerically Stable Softmax
Watch how subtraction of max, exponentiation, and normalization work on a small logit vector.
Complete Shape Trace
| Operation | Shape |
|---|---|
| scores input | (batch, T, T) |
| max(scores, dim=-1, keepdim=true) | (batch, T, 1) |
| scores - max | (batch, T, T) [broadcasted] |
| exp(shifted) | (batch, T, T) |
| sum(exp, dim=-1, keepdim=true) | (batch, T, 1) |
| exp / sum [broadcasted] | (batch, T, T) |
Key Insights
What Comes Next
Self-attention computes attention across all tokens. But in autoregressive models (language models, text generation), a token should never see future tokens. This is enforced by a causal mask that sets future scores to $-\infty$ before softmax. Next: Causal Attention.