Softmax overflow inside attention
Attention runs numbers through an exponential, which overflows a computer's number format alarmingly early - and the standard fix is one subtraction that changes no answer at all.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Attention has to raise a number to a power. That grows so fast it runs off the end of what a computer can store.
Think of the kitchen scale at a grocery shop. It goes up to five kilograms. Put a ten-kilo sack of rice on it and the needle slams to the top and stays there. The scale has not measured anything. It has run out of room.
Computers have exactly this limit. Every number has a largest value it can hold. Go past it and the answer becomes a special marker meaning "too big". Every calculation that touches it afterwards is ruined.
Attention walks up to that limit faster than almost anything else in a model.
Why attention gets close to the edge
To turn scores into shares, softmax raises a fixed number to the power of each score. Powers grow ferociously.
A score of 10 becomes about 22,000. A score of 20 becomes about 485 million. A score of 90, and the common number format used on graphics cards has already given up.
Ninety is not a large score. A model whose weights have drifted a little during training produces scores in the hundreds without anything visibly wrong.
The fix, and why it is free
Here is the trick, and it is genuinely elegant.
Shares are about proportions. Say three friends have 998, 999 and 1000 rupees. The shape of that comparison is the same as 0, 1 and 2. Subtracting the same amount from everyone changes nothing about how they compare.
So before raising anything to a power, subtract the largest score from every score. The largest becomes zero. Everything else becomes negative. Raising to a negative power gives a small positive number, and small numbers never overflow.
The shares that come out are identical. Not close — identical. The subtraction cancels out completely in the division that follows.
scores: 1000 1001 1002
raise to a power: overflow overflow overflow -> ruined
subtract the largest:
scores: -2 -1 0
raise to a power: 0.135 0.368 1.000
divide by the total: 0.090 0.245 0.665 -> correctEvery serious library does this. If you write softmax yourself and skip it, you have written a bug that will find you.
The other trap: half-size numbers
To go faster, models often use a smaller number format that holds half as many digits. It saves memory and time, and it is the normal way to run a model today.
Its ceiling is far lower. It gives up around sixty-five thousand.
Attention scores can reach that ceiling before the power step is even reached. That is a nastier failure, because there is no exponential involved — the multiplication itself ran out of room.
The standard answer is to compute the scores in the bigger format. Everything else stays in the smaller one. Every fast attention implementation does this internally.
A trap that catches almost everyone
A word is blocked by pushing its score to negative infinity. Its share then becomes zero.
Now imagine a row where every word is blocked. This happens with padded rows in a batch of sentences of different lengths.
The whole row is negative infinity. The largest score is negative infinity. Subtracting it gives negative infinity minus negative infinity, which is not a number. That "not a number" marker spreads through everything it touches. Your loss becomes not-a-number in a single step.
The fix is to block with a large negative number instead of true negative infinity. Or make sure padded rows never reach the loss.
Remember this
- Softmax raises numbers to a power, which overflows early.
- Subtracting the largest score first fixes it and changes no answer.
- A row where everything is blocked produces "not a number". That is the most common softmax bug in real code.
What to learn next
- QK normalisation — bounding the logits by construction instead of patching afterwards.
- Quantization in practice — the other place number formats decide whether a model works.
- Tensor creation and dtypes — the PyTorch side of choosing a format.
Developer — Code and libraries.
Setup
pip install numpyEvery failure mode, in one script
import numpy as np
np.set_printoptions(precision=4, suppress=True)
def naive_softmax(x):
e = np.exp(x) # nothing protects this line
return e / e.sum()
def stable_softmax(x):
e = np.exp(x - x.max()) # the largest input becomes exp(0) = 1
return e / e.sum()
logits = np.array([1000., 1001., 1002.])
with np.errstate(over="ignore", invalid="ignore"):
print("naive softmax([1000, 1001, 1002]) =", naive_softmax(logits))
print("stable softmax([1000, 1001, 1002]) =", stable_softmax(logits))
print("stable softmax([ 0, 1, 2]) =", stable_softmax(np.array([0., 1., 2.])))
print("-> shifting every logit by the same amount cannot change the answer")
print("\nwhere each number format gives up:")
print(" float32 largest finite value:", np.finfo(np.float32).max)
print(" float16 largest finite value:", np.finfo(np.float16).max)
with np.errstate(over="ignore"):
print(" float32: exp(88) =", np.exp(np.float32(88)), " exp(89) =", np.exp(np.float32(89)))
print(" float16: exp(11) =", np.exp(np.float16(11)), " exp(12) =", np.exp(np.float16(12)))
print("\nin float16 the logits can overflow before exp() is even reached:")
rng = np.random.default_rng(0)
d = 512
q = (rng.normal(size=(1, d)) * 50).astype(np.float16) # a head whose weights drifted large
k = (rng.normal(size=(4, d)) * 50).astype(np.float16)
with np.errstate(over="ignore"):
print(" q @ k.T computed in float16:", (q @ k.T)[0])
print(" q @ k.T computed in float32:", (q.astype(np.float32) @ k.astype(np.float32).T)[0])
print("\nthe mask bug nobody expects: a row where EVERY key is blocked")
with np.errstate(invalid="ignore"):
print(" stable_softmax([-inf, -inf, -inf]) =", stable_softmax(np.full(3, -np.inf)))
finite_floor = np.full(3, np.finfo(np.float32).min / 2, dtype=np.float32)
print(" same row with a large finite floor =", stable_softmax(finite_floor))naive softmax([1000, 1001, 1002]) = [nan nan nan] stable softmax([1000, 1001, 1002]) = [0.09 0.2447 0.6652] stable softmax([ 0, 1, 2]) = [0.09 0.2447 0.6652] -> shifting every logit by the same amount cannot change the answer where each number format gives up: float32 largest finite value: 3.4028235e+38 float16 largest finite value: 65500.0 float32: exp(88) = 1.6516363e+38 exp(89) = inf float16: exp(11) = 59870.0 exp(12) = inf in float16 the logits can overflow before exp() is even reached: q @ k.T computed in float16: [ 56992. -50336. 4640. inf] q @ k.T computed in float32: [ 56994.875 -50339.71 4641.293 102642.64 ] the mask bug nobody expects: a row where EVERY key is blocked stable_softmax([-inf, -inf, -inf]) = [nan nan nan] same row with a large finite floor = [0.3333 0.3333 0.3333]
Every line of that output is a lesson
Rows two and three are identical, to the last digit. [0.09, 0.2447, 0.6652] for inputs of 1000, 1001, 1002 and for 0, 1, 2. The shift is not an approximation. It is an exact identity that the implementation exploits.
float32 dies at 89, not at some huge number. The largest float32 is about 3.4 followed by 38 zeros. That sounds generous until you notice the exponential eats it in 89 steps. Attention logits reach 89 more easily than people expect.
float16 dies at 12. Twelve. That is the single most surprising number on this page. A model running in half precision has almost no headroom in the exponential at all.
The float16 matmul produced inf in the fourth slot. No exponential was involved. The dot product itself exceeded 65,500. Notice also 56992. against the float32 answer 56994.875. Even where it did not overflow, half precision lost the fractional part.
The all-blocked row gave [nan nan nan]. Subtracting negative infinity from negative infinity is undefined, and the result contaminates everything downstream. With a large finite floor instead, the same row gives a harmless uniform [0.3333, 0.3333, 0.3333].
What PyTorch already does for you
F.scaled_dot_product_attention handles the maximum subtraction internally. It accumulates the softmax statistics in float32, even for float16 or bfloat16 inputs. Written against PyTorch 2.5.1.
It does not rescue you from an all-blocked row. If you build masks yourself, that case is still yours to handle.
torch.nn.functional.log_softmax followed by nll_loss, or cross_entropy directly, are stable for the same reason. Never compute log(softmax(x)) as two separate steps. For a small share, the softmax underflows to zero and the logarithm gives negative infinity.
bfloat16 versus float16
Two half-size formats exist, and they fail differently.
| Format | Largest value | Precision |
|---|---|---|
| float16 | about 65,500 | about 3 decimal digits |
| bfloat16 | about 3.4e38 | about 2 decimal digits |
bfloat16 keeps float32's range and gives up precision instead. That is why it has become the default for training. Overflow is a hard failure. A little lost precision is usually absorbed by the optimiser. If you are debugging nan losses in float16, moving to bfloat16 is often the whole fix.
Finding the source of a nan
When a loss goes to nan, work backwards in this order:
torch.autograd.set_detect_anomaly(True) # slow, but names the offending op- Check for an all-masked row. Print
mask.all(dim=-1).any(). - Check the maximum absolute logit before softmax. If it is in the hundreds, the problem is upstream in the weights, not in softmax.
- Check for division by a norm that has become zero.
- Check for
logof a value that reached zero.
Attention logits are the most common source, and QK normalisation is the standard structural fix rather than a patch.
Common mistakes
Writing your own softmax without the maximum subtraction. It works in tests with small numbers and fails in production. There is no reason to hand-write it.
Using -float('inf') in a mask that will be combined with padding. Prefer a large finite negative value such as torch.finfo(dtype).min. Be aware that adding two of those can itself overflow to negative infinity. Setting the masked entries rather than adding to them avoids that.
Casting a mask built in float32 into a float16 forward pass. torch.finfo(torch.float32).min is far below the float16 range and becomes negative infinity on cast. Build masks in the dtype you will use.
Clamping logits as a permanent fix. Clipping does stop the nan. It also flattens gradients in the clipped region and hides a drift problem that will get worse. Use it to survive a run, then fix the cause.
Try it yourself
Take the float16 example and reduce the multiplier from 50 down to 30, then 20. Find the value where inf stops appearing. Then raise d from 512 to 2048 at the smaller multiplier and watch it come back. Score magnitude grows with both the weight scale and the head width. That is exactly why the scaling factor in attention depends on head width.
What to learn next
- QK normalisation — bounding the logits by construction instead of patching afterwards.
- Quantization in practice — the other place number formats decide whether a model works.
- Tensor creation and dtypes — the PyTorch side of choosing a format.
Researcher — Mathematics and papers.
The shift identity
For any constant $c$:
$$ \operatorname{softmax}(x)_i = \frac{e^{x_i}}{\sum_j e^{x_j}} = \frac{e^{x_i - c}}{\sum_j e^{x_j - c}} $$
Softmax is invariant to a uniform additive shift, so it is a function on $\mathbb{R}^n / \mathbb{R}\mathbf{1}$ rather than on $\mathbb{R}^n$. Taking $c = \max_j x_j$ places the largest exponent at $e^0 = 1$ and every other in $(0, 1]$. The numerator cannot overflow, and the denominator is at least 1. The remaining risk is underflow of very negative entries to zero, which is benign: those shares are genuinely negligible.
Overflow thresholds follow from the format:
$$ x_{\max}^{\text{fp32}} = \ln(3.40 \times 10^{38}) \approx 88.7, \qquad x_{\max}^{\text{fp16}} = \ln(65504) \approx 11.09 $$
Online softmax and why FlashAttention is exact
Milakov and Gimelshein (2018), arXiv:1805.02867, compute softmax in a single pass. They maintain a running maximum $m$ and running sum $\ell$. On seeing a new block with maximum $m'$:
$$ m_{\text{new}} = \max(m, m'), \qquad \ell_{\text{new}} = e^{m - m_{\text{new}}} \ell + e^{m' - m_{\text{new}}} \ell' $$
and accumulated outputs are rescaled by $e^{m - m_{\text{new}}}$. This is what makes tiled attention possible. It is why FlashAttention is bit-comparable to unfused attention rather than an approximation. The same shift identity applies blockwise.
Mixed precision in practice
Standard attention kernels keep $Q$, $K$, $V$ and the output in half precision. They accumulate $QK^\top$ and the softmax statistics in float32. Two independent reasons:
- Range. float16 logits overflow at 65,504, reachable by an ordinary dot product over a wide head with drifted weights.
- Accumulation error. Summing $d_k$ products in half precision loses low-order bits monotonically. Tensor-core matmuls accumulate in float32 by default for exactly this reason.
bfloat16 (Kalamkar et al., 2019, arXiv:1905.12322) keeps the float32 exponent field and truncates the mantissa to 7 bits. Range is preserved; precision is roughly 2 to 3 decimal digits. For training, range failures are catastrophic and precision failures are absorbed by the optimiser. That is why bfloat16 displaced float16 at large scale. Note that FP8 attention reintroduces range problems and generally requires per-tensor or per-block scaling.
Structural fixes for logit growth
Numerical hygiene stops the crash. It does not stop the underlying drift, which shows up as attention entropy collapsing toward zero.
- QK normalisation (Henry et al., 2020, arXiv:2010.04245) applies $L_2$ normalisation to $q$ and $k$ along the head dimension. The result is multiplied by a learned scale. Logits are then bounded by that scale. ViT-22B (Dehghani et al., 2023, arXiv:2302.05442) adopted it explicitly to fix divergence at scale. Several later open models followed.
- Logit soft-capping applies $c \tanh(s / c)$ to the pre-softmax scores, bounding them smoothly. Used in the Gemma 2 report. It is incompatible with some fused attention kernels, which is a practical reason it has not spread further.
- Reducing the entropy target directly. Zhai et al. (2023), arXiv:2303.06296, characterise entropy collapse as the precursor to divergence. They propose spectral reparameterisation of the query and key projections.
The mask floor, precisely
Using $-\infty$ for masked logits is correct whenever at least one entry per row is unmasked. When a row is fully masked, $\max_j x_j = -\infty$ and $x_i - \max_j x_j$ evaluates to $\text{NaN}$ under IEEE 754.
Substituting torch.finfo(dtype).min avoids that but introduces a second hazard: adding two such masks overflows to $-\infty$. The robust pattern is masked_fill rather than addition. Add an assertion that no row is fully masked. Exclude padded positions from the loss, so their outputs are never read.
Papers
- Milakov and Gimelshein, Online normalizer calculation for softmax, 2018 — arxiv.org/abs/1805.02867
- Micikevicius et al., Mixed Precision Training, 2018 — arxiv.org/abs/1710.03740
- Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training, 2019 — arxiv.org/abs/1905.12322
- Henry et al., Query-Key Normalization for Transformers, 2020 — arxiv.org/abs/2010.04245
- Zhai et al., Stabilizing Transformer Training by Preventing Attention Entropy Collapse, 2023 — arxiv.org/abs/2303.06296
What to learn next
- QK normalisation — bounding the logits by construction instead of patching afterwards.
- Quantization in practice — the other place number formats decide whether a model works.
- Tensor creation and dtypes — the PyTorch side of choosing a format.