FP16 attention partial-sum overflow in split attention kernels on M1 and M2
Parent: Mac local LLMs: MLX kernels, numerics and internals · Published reference · snapshot 2026-10-05
↓ Facts as markdownall context files
Pass 1 (`sdpa_vector_2pass_1`) holds its query, output accumulator, max and sum in float (`typedef float U`) and writes the block's running max and sum to float device buffers `maxs` and `sums`.
These notes link each claim to its source. A source may be a research report hosted on this site rather than the primary document. A published reference means the content is available; it does not certify independent review or accuracy.Read the editorial policy and follow the sources before relying on a claim.
Facts
- Pass 1 (`sdpa_vector_2pass_1`) holds its query, output accumulator, max and sum in float (`typedef float U`) and writes the block's running max and sum to float device buffers `maxs` and `sums`. [source]
- Its output accumulator is updated as `o[j] = o[j] * factor + exp_score * values[j]` and is never divided by the sum in pass 1, so the partial is an unnormalized weighted sum of values. [source]
- The partial is written with `out[i] = static_cast<T>(o[i])`, so it is rounded to fp16 for an fp16 model. [source]
- Pass 2 reads each partial as `T`, casts it to float, scales it by `exp(block_max - global_max)`, sums over blocks and divides by the total sum. The division happens in float after the fp16 rounding. [source]
- The 1-pass kernel differs: it divides by the sum in float before `static_cast<T>`, so its stored output is already a convex combination of values. [source]
- The steel (full, query length above 8) attention kernel keeps its score and output tiles in `AccumType = float`, so the prefill path has no fp16 partial. [source]
- Which chips reach 2-pass: chip class 'd' or 's' at 1024 keys or more, or any class with GQA at 4096 keys or more. [source]
- Blocks for classes other than 's' and 'd' are 64 when `n_simds` is 4 or more, else 32. [source]
- 2026-07-20: PR 3875's author notes in passing that the pass-2 partials offset arithmetic is 32-bit and would overflow once B x H x qL x blocks x D reaches 2^31. The PR does not change it. [source]
- 2026-08-05 to 2026-08-16: issue 4022 reports garbage decode from the first 2-pass step on an M3 Max; it was not reproduced elsewhere and was closed for inactivity. The report does not mention overflow. [source]
- Every exponential in a block is at most 1, so one block's fp16 partial is bounded by (keys in block) x (largest absolute value entry). For 8192 keys over 64 blocks that is 128 keys per block, so a block can exceed fp16's 65,504 only if its weighted values sum past that, for example 128 flat-weight keys with an average value magnitude above about 512. [source]
- Such value magnitudes are unusual but the Gemma 3 reports in existing dossiers show activations reaching about 800,000 in fp16, so the bound cannot be ruled out for those checkpoints. An overflowed partial becomes inf, and pass 2 then yields inf or NaN for that head and step. [source]
- The failure would depend on key length and block count, so it would start at the routing switch (1024 keys on Max and Ultra classes, 4096 keys with GQA on base and Pro classes), which is the signature issue 4022 describes for a different chip. [source]
- bfloat16 has fp32's exponent range, so the same cast cannot overflow in bf16; only fp16 models are exposed. [source]
- Existing m1-m2-software-emulated-bfloat16-and-fp16-conversion.md, citing Unsloth for Gemma 3, places fp16 overflow between decoder layers and "not inside attention/MLP". The kernel source leaves one fp16 rounding inside attention (the pass-1 partial) that could overflow, but no source shows it happening, so the two do not conflict on evidence. [source]
- `sdpa_vector_2pass_1` writes its unnormalized per-block output sum with `static_cast<T>`, so the partials buffer has the model dtype. [source]
- Pass 1 keeps block max and block sum in float buffers and its accumulator in float. [source]
- Pass 2 casts partials to float before weighting and dividing. [source]
- The 1-pass vector kernel divides by the sum before the cast to `T`. [source]
- The steel full-attention kernel uses a float accumulator type by default. [source]
- PR 3875's author noted that the pass-2 partials offset is 32-bit and overflows at B x H x qL x blocks x D of 2^31 or more. [source]
- Issue 4022 does not attribute its garbage output to fp16 overflow and was closed unresolved. [source]
- A pass-1 fp16 partial is bounded by keys-per-block times the largest value magnitude, because each exponential is at most 1. [source]
- No source found reports an fp16 partial-sum overflow in MLX's 2-pass attention on any chip, including M1 and M2. [source]
- bfloat16 models cannot overflow the pass-1 partial cast because bf16 shares fp32's exponent range. [source]
Children
- No children recorded.