M5 GPU TF32 numerics and batched attention divergence
Parent: Mac local LLMs: MLX kernels, numerics and internals · Published reference · snapshot 2026-10-05
↓ Facts as markdownall context files
Two independent mechanisms, not one. (1) fp16/bf16 masked attention takes the NAX attention kernel on gen-17; only an architecture override moves it, `MLX_ENABLE_TF32` does nothing to it. (2) fp32 GEMM takes the NAX GEMM at TF32 precision on gen-17; both `MLX_ENABLE_TF32=0` and a forced gen-16 ar...
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
- Two independent mechanisms, not one. (1) fp16/bf16 masked attention takes the NAX attention kernel on gen-17; only an architecture override moves it, `MLX_ENABLE_TF32` does nothing to it. (2) fp32 GEMM takes the NAX GEMM at TF32 precision on gen-17; both `MLX_ENABLE_TF32=0` and a forced gen-16 architecture move it. The thread started by attributing everything to the attention kernel and was corrected within three days. [source]
- Every Metal TF32 gate has the form `is_nax_available() && (env::enable_tf32() || dtype != float32)`. `enable_tf32()` reads `MLX_ENABLE_TF32` once with default 1 and latches, so it must be set before the first MLX op in the process. Because the third clause is true for any non-fp32 dtype, the flag can only switch fp32 off the NAX route. [source]
- `MLX_ENABLE_TF32` is inert on Metal before gen-17: M3 Max (`applegpu_g15s`) fp32 512^3 GEMM error was bit-identical with the flag on and off (2^-21.7; 2^-20.7 at 1024^3). On CUDA the flag is not gated by hardware generation at all. [source]
- `MLX_METAL_GPU_ARCH` overrides the architecture string that dispatch predicates read. Forcing `applegpu_g16s` on an M5 Max reproduced the M3 Max column cell for cell; `g14s` was byte-identical to `g16s`, so the responsible gate sits between gen-16 and gen-17. For fp32 rows this control is confounded because the override also disables TF32. [source]
- Masked-path batch-versus-single gap, M5 (gen-17), 32 seeds, H=16, Lreal=40, pad=8, median/max of max-abs difference: D=64 fp16 2^-11/2^-10, bf16 2^-8/2^-8, fp32 2^-11.7/2^-11.3; D=128 fp16 2^-11/2^-10, bf16 2^-8/2^-7, fp32 2^-11.7/2^-11.3. Same cells on M3 Max: D=64 fp16 2^-14.4/2^-11, bf16 0/2^-11, fp32 2^-22.1/2^-21.7; D=128 fp16 2^-12/2^-10, bf16 2^-11.4/2^-9, fp32 2^-22.0/2^-21.4. The unpadded control is exactly 0 on both. Two M5 variants (`g17g`, `g17s`) matched exactly. [source]
- The gap is heavy-tailed: an 8-seed sample made M3 Max bf16 look 512x better than M5 (2^-17 vs 2^-8); 32 seeds showed a median of 0 but a max of 2^-11 at D=64. Medians hid a 27-of-32 seed disagreement at D=96. Report distributions, not single medians. [source]
- Head-dim 96 initially looked like the clean case. On the v0.32.0 source quoted in the thread, 96 was outside the fused full-attention set (64, 80, 128), so it composed attention from ordinary matmuls and still responded to TF32 (27 of 32 seeds differ from native; `MLX_ENABLE_TF32=0` and forced g16s bit-identical to each other). On current main the NAX full-attention path takes D in (64, 96, 128, 256, 512), so the v0.32.0 head-dim explanation is version-specific. [source]
- Model-level effect: on mlx-lm tests/test_generate.py (commit 24b2b4e, mlx 0.32.0, M5 `g17g`) native gave 8 failed / 20 passed; `MLX_ENABLE_TF32=0` gave 28 passed, stable over three repeats, even though the model is a 4-bit fp16/bf16 model. This points to fp32 accumulation inside quantized matmul or other fp32 GEMMs, not the NAX attention kernel, as the dominant contributor to those eight failures. [source]
- Decode-token impact: Qwen1.5-0.5B-Chat-4bit, one decode step, batch of 4, max abs logprob difference 0.031-0.039 (about 1/32) with argmax matching in all four prompts. [source]
- Shape dependence: matvec shapes (M=1 or N=1) do not take the NAX GEMM route and stay exact fp32, so the same dtype and op has different precision depending on operand shapes. mlx-lm test_ssm passes on `out` (gemv-shaped at seq_len 1) and fails on `next_state` (outer-product GEMM). [source]
- 2026-03-10 (issue 3235): maintainer angeloskath: "TF32 is only supported by NAX. So this flag enables TF32 if NAX is available for this operation." The asker had read the gate as a speed switch. [source]
- 2026-05-12 (issue 3534): M5 and M5 Pro fp32 matmul precision regression "since 0.30.0" (0.29.4 fine). 256x256 randn vs torch CPU: M4 max abs diff 3.8e-5, M5 5.2e-2 (mean 9.8e-3), mlx 0.31.2, macOS 26.3.1. Maintainer zcbenz suggested `MLX_ENABLE_TF32=0`, which fixed it; reporter noted the variable must be set before any MLX computation. Closed as expected behavior. [source]
- 2026-07-17 (issue 3860): filed for CUDA (RTX 5070, sm_120, 512^3 fp32 GEMM rel error 2.9e-4 default vs 2.1e-7 with the flag off), retitled 2026-07-21 to "undocumented on both backends" after M5 Max data (outer product (64,1)@(1,128) max rel err 1.9e-3 vs 6.0e-8). Plan agreed: a doc line (commit 88fa4a0, 2026-07-23, "docs: document the TF32-class float32 default and MLX_ENABLE_TF32") and a one-time log line when fp32 GEMM actually engages reduced precision. The issue is closed. [source]
- 2026-07-21: mlx-lm PR 1595 pins `MLX_ENABLE_TF32=0` via `os.environ.setdefault` at import in tests/test_models.py so test_ssm (atol=rtol=1e-4, kernel vs reference) passes on M5 Max; full test_models 80 tests pass, 1 skipped. A thread comment says tests/test_generate.py was not covered by that pin. [source]
- 2026-07-23 to 2026-07-26 (issue 3897): batched attention divergence on M5 reported by mabaeyens (mlx-lm#1584 context), investigated by katlun-lgtm and PhilipJohnBasile. 2026-08-08: zcbenz closed it as won't fix: "it is an `MLX_ENABLE_TF32` behavior and for precision tests we should always turn it off." The fp16/bf16 NAX-attention half got no maintainer response in the thread. [source]
- Latching: setting `MLX_ENABLE_TF32` after the first matmul silently does nothing (confirmed on Metal and CUDA). In a server, set it in the launch environment, not in request-time code. [source]
- Op-level parity tests can pass while end-to-end decisions flip (CUDA report: 1.4-2.5% of argmax picks changed, about 9 dB PSNR lost) because the error enters through accumulated GEMMs. [source]
- The fp16/bf16 attention divergence cannot be turned off by `MLX_ENABLE_TF32`; the only user-level lever found is `MLX_METAL_GPU_ARCH` forcing a pre-gen-17 name, which also removes NAX GEMM speed. Its prefill cost was not measured in the thread. [source]
- A strict `rtol=1e-5` batch-equivalence assertion fails even in fp32 on gen-17, and even M3 Max bf16 reaches a max gap of 2^-9 at D=128, so such tolerances are fragile on any silicon at real head dims. [source]
- The Metal behavior depends on the macOS version: the issue's environment line first said macOS 15 (a typo for 26.5.2). On macOS before 26.2 the NAX path is unreachable and gen-16-like numerics appear. [source]
- Maintainer position (zcbenz): TF32 on by default is intended; tests must disable it. Reporter position (3860, 3897): the default silently changes numerics on both CUDA and gen-17 Metal, should default off or at least warn. Outcome: documented, not changed. [source]
- Whether the M5 fused-attention reduced-width accumulation is intended, and whether a maintainer will narrow it; no response in the cached thread. [source]
- Whether the one-time log line from the 3860 plan shipped in a release after 0.32.3. [source]
- Prefill and decode throughput cost of `MLX_ENABLE_TF32=0` on M5 for fp32 models; not measured in any source read. [source]
- Why TF32=0 alone fixes the quantized-model mlx-lm tests while TF32 does not touch fp16/bf16 attention; the exact fp32 op in the quantized path was not identified. [source]
- MLX defaults `MLX_ENABLE_TF32` to 1, read once into a static on first use. [source]
- On Metal, TF32 is only supported by NAX, so the flag enables TF32 only when NAX is available for the operation. [source]
- On M5 (gen-17) a 256x256 fp32 matmul against torch CPU had max abs diff 5.196e-2 versus 3.815e-5 on M4, and `MLX_ENABLE_TF32=0` fixed it. [source]
- The M5 fp32 precision change was first seen in mlx 0.30.0; 0.29.4 was fine. [source]
- On M3 Max (`applegpu_g15s`) the fp32 GEMM error is identical with `MLX_ENABLE_TF32` 0 and 1, at 512^3 and 1024^3. [source]
- On M5 base (`g17g`) fp32 512^3 GEMM rel error is 2^-10.4 native and 2^-19.8 with `MLX_ENABLE_TF32=0` or forced `g16s`; M5 Max `g17s` gave 2^-10.4 native and 2^-20.9 with either. [source]
- Forcing `MLX_METAL_GPU_ARCH=applegpu_g16s` on an M5 Max reproduced the M3 Max masked-attention error table cell for cell, and `g14s` gave byte-identical results to `g16s`. [source]
- With `MLX_ENABLE_TF32=0` on native gen-17, fp32 masked attention error falls from 2^-11.7 to 2^-22.1 (D=64) while fp16 and bf16 are unchanged. [source]
- M5 masked-path batch-versus-single gap at D=64 is 2^-11 (fp16), 2^-8 (bf16) and 2^-11.7 (fp32) median; M3 Max is 2^-14.4, 0 and 2^-22.1. [source]
- The unpadded batched-versus-single attention comparison is exactly 0 on both M5 and M3 Max at every head dim and dtype tested. [source]
- Two M5 variants (`applegpu_g17g` and `applegpu_g17s`) produced identical 32-seed tables. [source]
- mlx-lm tests/test_generate.py on an M5 fails 8 of 28 natively and passes 28 of 28 with `MLX_ENABLE_TF32=0`, stable over three runs (mlx-lm 24b2b4e, mlx 0.32.0). [source]
- Maintainer zcbenz closed issue 3897 on 2026-08-08 as won't fix, saying precision tests should always disable TF32. [source]
- Batched decode on M5 differs from single-sequence by about 0.031-0.039 max abs logprob (Qwen1.5-0.5B-Chat-4bit, batch 4) with argmax unchanged. [source]
- Every Metal TF32 gate in matmul, quantized matmul and SDPA has the form `is_nax_available() && (enable_tf32() || dtype != float32)`. [source]
- Issue 3860 shows fp32 outer-product and batched GEMM errors on M5 Max of 1.9e-3 and 4.3e-2 default versus 6.0e-8 and 2.8e-5 with `MLX_ENABLE_TF32=0`. [source]
- Matvec shapes (M=1 or N=1) do not take the NAX GEMM route and stay exact fp32. [source]
- The `MLX_ENABLE_TF32` variable latches on first use, so setting it after the first matmul does nothing. [source]
- The documentation commit "docs: document the TF32-class float32 default and MLX_ENABLE_TF32" (88fa4a0) was linked to issue 3860 on 2026-07-23. [source]
- mlx-lm PR 1595 pins `MLX_ENABLE_TF32=0` with `os.environ.setdefault` at import so test_ssm passes on M5 Max, and states M1-M4 are unaffected (no NAX, fp32 GEMM exact). [source]
- In test_ssm the `out` comparison passes because seq_len=1 contractions are gemv-shaped, while the `next_state` outer-product GEMM carries about 2e-3 relative TF32 error. [source]
- On current main, full attention takes the NAX kernel for head dims 64, 96, 128, 256, 512 when `is_nax_available()` and (TF32 enabled or dtype is not fp32). [source]
- Plain batched attention without padding is bit-identical to single-sequence on M5; only the padded and masked path used by batched generation diverges. [source]
- For server operators on gen-17: set `MLX_ENABLE_TF32=0` in the launch environment before any MLX op if fp32 reproducibility matters; it does not remove the fp16/bf16 attention gap. [source]
Corrections and disagreements
- CONTRADICTS kv-cache-quantization-tradeoffs-on-apple-gpus.md: it attributes the eight failures to TF32 in fp32 GEMM and lists "31" tests; the issue author's counts are 8 of 28 (28 pass with TF32 off). It also folds the batched-attention divergence into one cause; the thread separates a NAX-attention cause (fp16/bf16, not TF32-controlled) from the TF32 GEMM cause (fp32 and apparently the quantized-matmul path). The count may differ because the suite changed. [source]
Children
- No children recorded.