Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

13. Reductions, softmax and floating point

In this chapter

  • Parallel reductions: trees, shared memory and warp shuffles, the pattern behind softmax, norms and losses.
  • Floating-point formats from FP32 to FP4: what they can represent, and what they silently can't.
  • Numerically stable softmax and log-sum-exp, and why accumulation precision matters.
  • How to compare two implementations honestly: tolerances, aggregate errors, and the right reference.

You will build

stable_softmax in engine/numerics.py and, on a GPU, warp- and block-level reduction kernels, ending with an RMSNorm kernel.

Time: 5-7 hours. GPU: optional.

One answer from many values

A reduction combines many values into one: a sum, a maximum, a mean. Transformers are full of them. Softmax needs a maximum and a sum per row, LayerNorm and RMSNorm need a mean of squares per token, and the loss averages over tokens. Unlike elementwise operations, reductions need threads to combine their results, which raises two questions: how to do it in parallel, and how much numerical precision the combination loses.

Parallel reduction: a tree

A serial sum does $N - 1$ additions one after another. A tree adds independent pairs in parallel: $N/2$ additions, then $N/4$, and so on. It’s still $N-1$ additions in total, but only $\log_2 N$ sequential steps. For 8 values, that’s 3 steps instead of 7.

On a GPU, a block reduces its share of the data in two levels:

  1. Within a warp, threads exchange register values directly with shuffle instructions. __shfl_down_sync(mask, v, offset) gives each thread the value held by the thread offset lanes above it. Five rounds (offsets 16, 8, 4, 2, 1) leave the warp’s total in lane 0, with no shared memory and no barriers.
  2. Across warps, each warp’s lane 0 writes its partial sum to shared memory, the block synchronizes, and the first warp reduces those partials the same way.
// Sum within a warp: each step halves the number of live partial sums, using register
// exchanges instead of shared memory. After 5 steps lane 0 holds the warp's total.
__device__ float warp_sum(float v) {
    for (int offset = 16; offset > 0; offset >>= 1) v += __shfl_down_sync(0xffffffff, v, offset);
    return v;
}

// Block-wide sum: every warp reduces, warp leaders publish to shared memory, warp 0 finishes.
__device__ float block_sum(float v) {
    __shared__ float partial[32];
    int lane = threadIdx.x % 32, warp = threadIdx.x / 32;
    v = warp_sum(v);
    if (lane == 0) partial[warp] = v;
    __syncthreads();
    v = (threadIdx.x < blockDim.x / 32) ? partial[lane] : 0.f;
    if (warp == 0) v = warp_sum(v);
    __syncthreads();                             // partial[] may be reused by the caller
    return v;                                    // valid in thread 0
}

// One block per row; each thread strides over the row (thread coarsening), then reduces.
__global__ void row_sum_kernel(const float* x, float* out, int width) {
    const float* row = x + int64_t(blockIdx.x) * width;
    float v = 0.f;
    for (int i = threadIdx.x; i < width; i += blockDim.x) v += row[i];
    v = block_sum(v);
    if (threadIdx.x == 0) out[blockIdx.x] = v;
}

row_sum_kernel also uses thread coarsening: each thread first sums many elements serially in a register (for (i = tid; i < width; i += blockDim.x)), and only then does the block reduce 256 partials. Reading consecutive elements per iteration keeps the loads coalesced. Two details are worth remembering:

  • Identity values. Threads with no data contribute the operation’s identity: 0 for a sum, but $-\infty$ for a max. Contributing 0 to a max of all-negative values would be wrong.
  • Barriers across blocks don’t exist. If a reduction spans several blocks, write per-block partial results to global memory and reduce them in a second kernel, or use atomics.

RMSNorm: a reduction plus a broadcast

RMSNorm (Chapter 17) needs one reduction per token (the mean of squares), then the result must reach every thread so each can scale its elements. Thread 0 writes the statistic to shared memory, the block synchronizes, and everyone reads it:

__global__ void rmsnorm_kernel(const float* x, const float* w, float* y, int width, float eps) {
    const float* row = x + int64_t(blockIdx.x) * width;
    float* out = y + int64_t(blockIdx.x) * width;
    __shared__ float inv_rms;
    float sq = 0.f;
    for (int i = threadIdx.x; i < width; i += blockDim.x) sq += row[i] * row[i];
    sq = block_sum(sq);
    if (threadIdx.x == 0) inv_rms = rsqrtf(sq / width + eps);
    __syncthreads();                             // broadcast the statistic to every thread
    for (int i = threadIdx.x; i < width; i += blockDim.x) out[i] = row[i] * inv_rms * w[i];
}

One block per token row, reading the row twice: once for the statistic, once to normalize. A fused Triton version (Chapter 14) also adds the residual connection in the same pass.

Floating point: what your numbers can hold

A floating-point number is $\pm, 1.m \times 2^{e}$: a sign bit, an exponent field, and a mantissa (fraction) field. The exponent bits set the range, and the mantissa bits set the precision, meaning the gap between neighbouring representable numbers. Every format makes a different trade:

formatsign / exponent / mantissa bitslargest valuegap just above 1.0typical use
FP321 / 8 / 233.4 × 10³⁸1.2 × 10⁻⁷accumulation, optimizer state
BF161 / 8 / 73.4 × 10³⁸0.0078weights and activations
FP161 / 5 / 1065,5040.00098older GPUs, some kernels
FP8 E4M31 / 4 / 34480.125weights, activations (scaled)
FP8 E5M21 / 5 / 257,3440.25gradients (scaled)
FP4 E2M11 / 2 / 160.5weights, with per-block scales

BF16 keeps FP32’s 8-bit exponent, so it has the same range and never overflows where FP32 wouldn’t. It pays with precision: only about 2-3 significant decimal digits. That’s why it dominates training and inference. FP16 has more precision but overflows above 65,504, a real problem for activations. FP8 and FP4 are so coarse that they only work with scale factors stored alongside groups of values (Chapter 20).

Explore: floating-point bits

Type a number and see its sign, exponent and mantissa bits in each format, the value actually stored, and the rounding error.

Rounding makes order matter

In exact arithmetic, $(10^8 + 1) - 10^8 = 1$. In FP32, the gap between representable numbers near $10^8$ is 8, so $10^8 + 1$ rounds back to $10^8$ and the result is 0. Computing $(10^8 - 10^8) + 1$ instead gives 1. Floating-point addition is not associative: a tree reduction, a fused kernel and a serial loop can all give different answers, all legitimately.

Precision of the accumulator matters even more. Summing 0.1 ten thousand times should give 1,000:

BF16 accumulator, sequential:  32.0
FP32 accumulator:              1000.98    (0.1 itself isn't exact in BF16)

The BF16 sum is stuck at 32, where the gap between representable numbers is 0.25: adding 0.1 rounds back down every time. This is why every serious kernel stores in low precision but accumulates in FP32: matmuls on tensor cores, reductions in norms, the softmax denominator. Your attention and norm code converts to FP32 (.float()) for its reductions for exactly this reason.

Stable softmax

$e^{1000}$ overflows every format, so a naive softmax of [1000, 999, 998] returns [nan, nan, nan]. Since softmax is unchanged by subtracting a constant from every input,

$$ \operatorname{softmax}(x)_i = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}}, \qquad m = \max_j x_j , $$

subtracting the row maximum makes the largest exponent $e^0 = 1$ and everything else smaller. The result is $[0.6652, 0.2447, 0.0900]$, exactly as if the inputs had been $[2, 1, 0]$.

def stable_softmax(scores, dim=-1):
    """Softmax that subtracts the row maximum first.  (Your engine: Chapter 13)

    Computes in FP32 and returns the input dtype. Each row needs at least one finite score.
    """
    x = scores.float()
    shifted = x - x.amax(dim=dim, keepdim=True)
    numerators = shifted.exp()
    return (numerators / numerators.sum(dim=dim, keepdim=True)).to(scores.dtype)
inline void softmax(float* x, size_t n) {
    float m = *std::max_element(x, x + n), sum = 0;
    for (size_t i = 0; i < n; ++i) sum += (x[i] = std::exp(x[i] - m));
    for (size_t i = 0; i < n; ++i) x[i] /= sum;
}
#![allow(unused)]
fn main() {
/// Stable softmax in place: subtract the maximum so the largest exponent is exp(0) = 1.
pub fn softmax(x: &mut [f32]) {
    let max = x.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
    let mut sum = 0.0;
    for v in x.iter_mut() {
        *v = (*v - max).exp();
        sum += *v;
    }
    for v in x.iter_mut() {
        *v /= sum;
    }
}
}

The same trick gives a stable log-sum-exp, $\log\sum_j e^{x_j} = m + \log\sum_j e^{x_j - m}$, which is how cross-entropy is computed without ever forming probabilities (Chapter 7). Keep the softmax’s edge case in mind: if every entry of a row is $-\infty$ (a fully masked attention row), the max is $-\infty$ and the result is NaN. Your mask logic must guarantee at least one visible key, or define what an empty row means.

In Chapter 15, you’ll extend this to compute softmax in tiles, carrying the running maximum forward. That’s the online softmax at the heart of FlashAttention.

Comparing implementations honestly

Because different correct implementations round differently, “equal” means “close enough”, and you must decide how close before you look at the results. The standard test combines absolute and relative tolerances:

$$ |a - r| \le \text{atol} + \text{rtol}\cdot|r| . $$

The absolute term handles values near zero; the relative term handles scale. Good error reports also include aggregate measures, because a maximum alone can hide a systematic drift:

def compare(actual, expected):
    """An error report, not a pass/fail verdict. Choose tolerances per dtype and operation."""
    difference = (actual.float() - expected.float()).abs()
    reference = expected.float()
    return {
        "max_abs": difference.max().item() if difference.numel() else 0.0,
        "mean_abs": difference.mean().item() if difference.numel() else 0.0,
        "relative_l2": (
            torch.linalg.vector_norm(difference)
            / torch.linalg.vector_norm(reference).clamp_min(1e-12)
        ).item(),
    }

Which tolerances? They depend on dtype, operation and magnitude. A BF16 result can’t be closer than about 0.4% relative to anything, so atol=1e-5 on BF16 logits is meaningless. A useful discipline is to make three comparisons:

  1. Trusted FP32 vs your FP32: tests your implementation. Expect ~1e-5 or better.
  2. Trusted BF16 vs trusted FP32: measures what the precision itself costs.
  3. Your BF16 vs trusted BF16: should be about as close as (2), not much worse.

If (3) is much worse than (2), you have a bug or a bad accumulation choice. And whatever the tolerance, rounding never excuses a semantic error: a wrong mask or transpose produces differences far larger than any rounding.

Warning

On NVIDIA GPUs, PyTorch may run “FP32” matmuls with TF32 tensor cores (10-bit mantissa) when torch.backends.cuda.matmul.allow_tf32 is on, giving about 1e-3 relative error. Set the flag explicitly whenever you compare FP32 results.

Note

Some GPU reductions are nondeterministic: using atomic additions or split-K matmuls, the order of additions changes run to run, and so do the last bits of the result. GPU Mode L9 (nondeterminism.py) demonstrates it. Serving engines that promise bitwise-reproducible outputs have to avoid these kernels.

Build it

Engine milestone 13: stable numerics. Implement stable_softmax in engine/numerics.py (logsumexp, compare and float_bits are provided). On a GPU, implement warp_sum, block_sum, row_sum_kernel and rmsnorm_kernel in engine/kernels/cuda_ops.cu.

pytest tests/test_ch13_numerics.py
pytest tests/test_ch11_cuda.py -k reductions     # GPU

Stretch exercises

  1. ★ Reproduce the BF16 summation failure, then implement Kahan summation in BF16 (carry a compensation term) and see how close it gets to 1,000. Where: experiments/ch13.py (create it); add a Kahan helper to engine/numerics.py.
  2. ★★ Write a block reduction for the maximum and check it on rows of all-negative numbers. What breaks if invalid lanes contribute 0? Where: add max-reduction helpers and a row-max kernel, launcher and binding in engine/kernels/cuda_ops.cu.
  3. ★★ Run Qwen3 or your GPT in FP32 and in BF16 on the same input, and report max-abs and relative-L2 error of the logits, and whether greedy tokens agree over 50 steps. Where: experiments/ch13.py (create it), importing engine.numerics.compare.
  4. ★★★ Write a single-pass softmax kernel in CUDA: one block per row, a max reduction, a sum reduction, then normalize, with the row held in registers. Compare its bandwidth with torch.softmax. Where: add a row-softmax kernel, launcher and binding in engine/kernels/cuda_ops.cu.

Check your understanding

  1. Why does subtracting the row maximum leave softmax unchanged?
  2. Why does BF16 have FP32’s range but far less precision?
  3. Why should a BF16 kernel accumulate sums in FP32?
  4. Why report aggregate error in addition to the maximum error?
  5. What identity value should inactive threads contribute to a max reduction?

Going deeper

  • PMPP Chapter 10 (reductions: trees, divergence, warp shuffles, coarsening, pp. 225-248) and Appendix A (floating-point considerations).
  • GPU Mode L9 (Mark Saroufim, reductions, with runnable kernels and the nondeterminism demo) and L84 (numerics and AI).
  • David Goldberg, What Every Computer Scientist Should Know About Floating-Point Arithmetic (1991); Micikevicius et al., FP8 Formats for Deep Learning (2022).