PidokuInfra

Softmax and Numerical Stability

Foundations Intermediate 1h 15m Difficulty 3/5 Topic 07 of 12

Prerequisites 01, 03

The “online softmax” trick in section 6 is the mathematical core of FlashAttention. If you understand it here, Section VII.10 will be easy.


1. What is it?#

Softmax turns a vector of arbitrary real numbers into a probability distribution:

softmax(x)_i = exp(x_i) / sum_j exp(x_j)

Every output is in (0,1) and they sum to 1. It appears twice in every LLM:

  1. Inside attention, over the key positions.
  2. At the output, over the vocabulary, to sample the next token.

2. Why does it exist?#

Because you have scores and you need probabilities, and you want three properties:

  1. All positive — exp guarantees it.
  2. Sum to 1 — the normalization guarantees it.
  3. Monotonic and differentiable — bigger score → bigger probability, smoothly.

The exponential also means differences matter multiplicatively: a score gap of 2 produces a probability ratio of e² ≈ 7.4, regardless of the absolute values. That “relative gaps decide” property is why it’s called softmax — it’s a smooth approximation to “pick the max.”


3. Simple analogy#

Election results from raw enthusiasm scores.

Three candidates receive enthusiasm scores 8, 6, 2. You need vote shares. You could normalize linearly (8/16, 6/16, 2/16), but that treats a score of 0 as no support and can’t handle negatives. Softmax exponentiates first: e⁸, e⁶, e² → 88%, 12%, 0.03%. The exponential sharpens the difference — the leader wins decisively.

Temperature (file V.13) is the dial that controls how sharp: divide scores by T before exponentiating. High T flattens (more random), low T sharpens (more deterministic).


4. Tiny example#

Go
package main

import (
	"fmt"
	"math"
)

func softmaxNaive(x []float64) []float64 {
	out, sum := make([]float64, len(x)), 0.0
	for i, v := range x {
		out[i] = math.Exp(v) // [7.389, 2.718, 1.105]
		sum += out[i]
	}
	for i := range out {
		out[i] /= sum
	}
	return out
}

func main() {
	fmt.Printf("%.3f\n", softmaxNaive([]float64{2.0, 1.0, 0.1})) // [0.659 0.242 0.099], sums to 1 ✓
}

Now break it:

Go
fmt.Println(softmaxNaive([]float64{1000, 999, 998}))
// math.Exp(1000) = +Inf, and Inf/Inf = NaN:
// [NaN NaN NaN]        ← OVERFLOW

The correct answer is the same as for [2,1,0]: [0.665, 0.245, 0.090], because softmax only depends on differences. But naive computation produces NaN.

The fix — subtract the max:

Go
func softmaxStable(x []float64) []float64 {
	m := slices.Max(x) // subtract the max: the largest becomes 0, so max(exp) = 1
	out, sum := make([]float64, len(x)), 0.0
	for i, v := range x {
		out[i] = math.Exp(v - m)
		sum += out[i]
	}
	for i := range out {
		out[i] /= sum
	}
	return out
}

// softmaxStable([]float64{1000, 999, 998})
// [0.665 0.245 0.090]        ✓

Why this is exactly correct, not an approximation:

exp(x_i - c)/Σ exp(x_j - c) = [exp(x_i)·exp(-c)] / [exp(-c)·Σ exp(x_j)]
                             = exp(x_i)/Σ exp(x_j)

The exp(-c) cancels. Any c works; c = max(x) is chosen because it guarantees the largest exponent is exp(0) = 1, so nothing overflows.


5. Technical explanation#

Range limits by precision#

FP32:  max ≈ 3.4e38    → exp overflows around x > 88
FP16:  max ≈ 65,504    → exp overflows around x > 11.09
BF16:  max ≈ 3.4e38    → same range as FP32 (but 8 bits of mantissa)

FP16 overflows at x > 11. Attention logits routinely exceed that at long context. This is why:

  • Attention softmax is computed in FP32 even in FP16/BF16 models.
  • BF16 is preferred over FP16 for inference: same range as FP32 avoids a whole class of overflow bugs, at the cost of precision you mostly don’t need.

Underflow#

exp(-100) ≈ 3.7e-44   → subnormal in FP32, zero in FP16

Underflow to zero is usually harmless in softmax (that token had negligible probability anyway). It matters when you then take log: log(0) = -inf. Compute log-probabilities with log-sum-exp instead:

Go
func logSoftmax(x []float64) []float64 {
	c, sum := slices.Max(x), 0.0
	for _, v := range x {
		sum += math.Exp(v - c)
	}
	out := make([]float64, len(x))
	for i, v := range x {
		out[i] = v - c - math.Log(sum)
	}
	return out
}

Always use log_softmax rather than log(softmax(x)).

Attention softmax specifics#

scores = (q @ k.T) / sqrt(d_head)     # (S_q, S_k)
scores = scores + causal_mask          # masked positions get -inf (or -1e9)
attn = softmax(scores, dim=-1)         # over KEYS
out = attn @ v

Three subtleties:

1. The mask. Use -inf where possible so exp(-inf) = 0 exactly. -1e9 in FP16 overflows to -inf anyway; -1e4 does not fully mask. A row that is entirely masked gives 0/0 = NaN — this happens with certain padding schemes, and is a classic bug.

2. The scale. 1/sqrt(d_head) keeps the logits’ variance ~1 regardless of head dimension (Section III.01). Without it, softmax saturates and gradients vanish (training) / attention becomes one-hot (inference).

3. Which axis. Softmax is over the key dimension (last axis). Getting the axis wrong produces a valid-shaped tensor and nonsense.

Online (streaming) softmax — the FlashAttention core#

The problem: computing softmax normally requires the whole score vector in memory to find the max and the sum. For S = 128,000 keys, that vector is large, and you must materialize it.

The insight: you can compute softmax incrementally, updating a running max and a running sum as blocks arrive.

Process keys in blocks. Maintain:
    m = running maximum
    l = running sum of exp(x - m)
    o = running weighted output

For a new block with max m_new and values x:
    m'   = max(m, m_new)
    l'   = l · exp(m - m') + Σ exp(x - m')
    o'   = o · exp(m - m') + Σ exp(x - m') · v

The exp(m - m') factor rescales everything computed so far to the new maximum. When the last block is processed, divide o by l and you have the exact same answer as full softmax.

Concretely, two blocks:
  Block 1: x = [2, 1]     → m=2, l = e⁰ + e⁻¹ = 1.368
  Block 2: x = [5, 0]     → m_new = 5, m' = 5
           rescale: l = 1.368 · e^(2-5) = 1.368 · 0.0498 = 0.0681
           add:     l += e⁰ + e⁻⁵ = 1.0067  →  l = 1.0748
  Full softmax denominator for [2,1,5,0] with max 5:
           e⁻³ + e⁻⁴ + e⁰ + e⁻⁵ = 0.0498+0.0183+1+0.0067 = 1.0748   ✓ EXACT

This is not an approximation. It is exact, and it means attention never needs the full (S, S) matrix in memory. That is FlashAttention (Section VII.10), and it is why attention at 128k context is feasible at all.


6. Under the hood#

A naive softmax kernel makes three passes over the data:

pass 1: find max
pass 2: compute exp and sum
pass 3: divide
→ 3 reads + 1 write of the full tensor

A fused “online” kernel makes one pass, keeping running max and sum in registers:

pass 1: single pass, running max and sum, then normalize
→ 1 read + 1 write (or 2 reads if a second pass over cached values is needed)

For the LM head softmax over V=128256 at batch 64, that is the difference between reading ~65 MB and ~16 MB per step. Modest but real.

For attention, the difference is between materializing (B,h,S,S) in HBM and never materializing it. That is the difference between OOM and working.


7. Performance implications#

  • Softmax is memory-bound. ~5 FLOPs per element, 4-8 bytes moved. Intensity ~1.
  • Fusing softmax into attention is the single largest attention optimization.
  • FP32 accumulation for the sum costs nothing meaningful (the sum is one value per row) and buys stability.
  • The LM head softmax at large vocabulary is nontrivial: (B, 128256) reduction per step. Fuse it with sampling.

8. Production implications#

  • Use BF16 over FP16 for inference where possible. The range advantage eliminates a class of overflow bugs, especially at long context.
  • Always compute attention softmax with FP32 accumulation, even in a BF16 model.
  • Watch for NaN. A NaN in production usually traces to: a fully-masked attention row, an FP16 overflow in attention logits, or a division by a zero-sum. Add a debug mode that checks for NaN after each layer.
  • log_softmax for logprobs. Never log(softmax(x)).
  • Temperature 0 should not be implemented as softmax with T=0 (division by zero). Special- case it as argmax.

9. Common mistakes#

Naive softmax without max subtraction. Overflows immediately in FP16.

log(softmax(x)) instead of log_softmax(x). Underflow → -inf → NaN downstream.

Wrong axis. Softmax over the query dimension instead of the key dimension. Shape is valid; results are garbage.

Using -1e9 as a mask in FP16. It overflows to -inf (fine) but -1e4 does not mask adequately, and the choice interacts with precision.

Fully-masked rows. All keys masked → sum = 0 → NaN. Happens with certain padding or sliding-window edge cases.

Forgetting the 1/sqrt(d_head) scale. Attention collapses to near-one-hot.

Computing softmax in FP16 to “save memory.” The memory saving is negligible; the risk is not.


10. Hands-on exercise#

A. Break and fix. Implement naive and stable softmax. Find inputs that make the naive version produce inf/NaN in FP32 and in FP16. Verify the stable version handles them.

B. Implement online softmax. Write a function that computes softmax over a long vector by processing it in blocks of 128, maintaining running max and sum. Verify it matches scipy.special.softmax to floating-point precision. Keep this code — it’s the heart of Project 04.

C. Extend to attention. Extend B to compute softmax(QK^T/√d) @ V block by block, never materializing the full score matrix. Verify against the naive implementation. Measure peak memory for both at S = 4096. You have now implemented FlashAttention’s core algorithm.

D. Range experiment. For d_head ∈ {64, 128, 256}, generate random q, k ~ N(0,1) and measure the distribution of q·k with and without the 1/sqrt(d) scale. At what d_head does unscaled attention overflow FP16?

E. NaN hunt. Construct an attention input with a fully-masked row. Observe the NaN. Fix it.


11. Interview questions#

  1. Why does softmax subtract the maximum? Prove it doesn’t change the result.
  2. At what value does exp overflow in FP16? Why does that matter for attention?
  3. What is online softmax and why does FlashAttention need it?
  4. Why is BF16 preferred over FP16 for LLM inference?
  5. Your model produces NaN after 3,000 tokens. Give four hypotheses.
  6. Why is log_softmax preferred to log(softmax(x))?
  7. Why does attention divide by sqrt(d_head)?

12. Further reading#

  • [ESTABLISHED] Milakov & Gimelshein, “Online normalizer calculation for softmax” (2018) — the original online softmax paper
  • [ESTABLISHED] Dao et al., “FlashAttention” (2022) §3.1
  • [FUNDAMENTAL] Goldberg, “What Every Computer Scientist Should Know About Floating-Point Arithmetic” (1991)
  • Next: 08 — Attention from first principles

↑↓ navigate↵ openesc close