Skip to content

fix(reductions): accumulate half-precision reductions in f32 - #1718

Merged
inureyes merged 2 commits into
mainfrom
fix/half-precision-reduction-nan
Sep 8, 2026
Merged

inureyes merged 2 commits into
mainfrom
fix/half-precision-reduction-nan

Conversation

@inureyes

@inureyes inureyes commented Sep 8, 2026

Copy link
Copy Markdown
Member

Closes #997.

text_only_forward_produces_finite_logits has failed intermittently since 2026-08-02. The fault is not in the model: reading the logits row back to the host found every element finite on 60 of 60 runs while the device reduction over the same row reported non-finite on 28, with no overlap. max over a bfloat16 array returns NaN for a finite input on M5 Max, about one call in six.

Measured one variant per process over 150 forwards of granite-4.0-3b-vision-4bit: 23 non-finite for the plain reduction, 31 reshaped to 1-D, 20 made contiguous, 24 reduced over an axis, 0 cast to f32. Element width is the trigger; contiguity and the flat-versus-axis choice are not. sum_all fails the same way at 28 of 150, so the whole reduction family promotes to f32 and restores the input dtype.

What this cost to find, and what stops the next one

Four hypotheses were eliminated by measurement before the right one: the Metal 4 attention route (MLXCEL_METAL4_ATTENTION=0 changed nothing), the grouped GQA decode path, the granitemoehybrid backbone (two text checkpoints on it are 0 of 40), and repeated make_caches. The M1 Ultra sees none of it.

Two traps produced wrong intermediate answers and are recorded in the code:

Timing several variants inside one loop hides the fault, because each eval synchronises the next one. The 1-D reshape read 0 of 60 measured that way and 123 of 150 as the only variant in its process, which would have been reported as a fix.

cargo build reported a 0.12s no-op after the bridge .cpp changed, so the first verification ran against a binary without the fix. Touching build.rs forced the 2m50s rebuild that actually contained it.

The guard

c07826df removed the M5 per-mixer eval from the Mamba2 hybrids and said mamba2_hybrid_decode_is_finite guarded the regression. It did not exist: the name appeared only in two comments pointing at each other. It exists now, and it took two corrections to make it discriminate. Sixty iterations passed without the fix, so it guarded nothing; and running the tiny checkpoint first in the same function stopped the fault appearing at all, so each checkpoint gets its own test. With the fix reverted it fails 4 of 4; with it applied it passes 4 of 4.

Also fixes handoff_timeout_check, which failed about one gate run in three on a clock-granularity race its own comment had flagged.

Gate: fmt, clippy --workspace --all-targets, and four consecutive workspace runs at 10534 passing with no failures.

A flat `max` over a bfloat16 array returns NaN for a finite input on M5 Max, on roughly one call in six. It is the reduction and not the data: reading the same row back to the host found every element finite on 60 of 60 runs while the device reduction reported non-finite on 28 of them, with no overlap between the two.

Casting to f32 first is the only thing that fixes it. Measured one variant per process over 150 forwards of `granite-4.0-3b-vision-4bit`, whose logits are bfloat16: 23 non-finite for the plain reduction, 31 reshaped to 1-D, 20 made contiguous, 24 reduced over an axis, and 0 cast to f32. Contiguity and the flat-versus-axis choice are not the trigger; the element width is. `sum_all` fails the same way at 28 of 150, so `max_all`, `min_all`, `sum_all`, `mean_all` and the four axis forms all promote now and restore the input dtype on the way out. Promotion is exact for max and min, whose result is one of the inputs, and strictly better for sum and mean.

Two measurement traps are worth recording. Timing several variants inside one loop hides the fault, because each `eval` synchronises the next variant: the 1-D reshape read 0 of 60 that way against 123 of 150 as the only variant in its process. And `cargo build` reported a 0.12s no-op after the bridge changed, so the first verification ran against a binary without the fix; touching `build.rs` forced the 2m50s rebuild that actually contained it.

`llama4.rs` already cast to f32 before these calls, which reads as an earlier encounter with the same fault. Its comment now says what the cast is and is not for.
`c07826df` dropped the M5 per-mixer `eval` from the Mamba2 hybrids and recorded that `mamba2_hybrid_decode_is_finite` guarded the regression it had been protecting against. No such test existed. The name appeared twice in the tree, in the `granitemoehybrid.rs` and `falcon_h1.rs` comments, each pointing at the other and neither at a definition, so the safety claim rested on something that was never written.

The test now exists and discriminates, which took two corrections. Sixty iterations passed without the fix under test, so it guarded nothing: raised to 250 after that. Exercising the tiny checkpoint before the vision one in the same function also stopped the fault appearing at all, so each checkpoint gets its own test function. With those, the reduction fix reverted, it fails 4 runs of 4; with it applied, it passes 4 of 4.

Also fixes a timing race in `handoff_timeout_check`, which failed about one gate run in three. `is_timed_out` is a strict `elapsed() > timeout`, so a zero timeout is not satisfied until the clock has advanced a tick, and the test asserted it immediately. Its own comment had flagged the race. It now waits for the tick instead of assuming it, keeping the case the test is about.
@inureyes
inureyes merged commit 0accedd into main Sep 8, 2026
13 checks passed
inureyes added a commit that referenced this pull request Sep 9, 2026
`repo_model_dir` looked only under `models/<name>`, and every caller treats a missing directory as "skip this test". The store was later consolidated into `models/mlx/`, with checkpoints over 120GB in `models/mlx-big/`, so the lookup stopped resolving and the tests went silently green instead of failing. On this M1 Ultra tree none of the 30 distinct names used across `tests/` resolved under `models/<name>`, so every real-model integration test in the repository was a no-op.

`tests/mamba2_hybrid_finite.rs` was one of them: the NaN guard written in #1718 ran in 0.00s and asserted nothing. With the roots fixed it runs 250 forwards per checkpoint in 10.64s.

`MLXCEL_REQUIRE_MODELS=1` turns a name that resolves nowhere into a panic naming every root tried, so a local gate run cannot pass by skipping everything. CI has no checkpoints and leaves it unset, keeping the skip there.

Fifteen of the 30 names now resolve. The other 15 are renames from the same consolidation (`llama-3.1-8b-4bit` is `meta-llama-3.1-8b-instruct-4bit`, `qwen2.5-7b-4bit` is `qwen2.5-7b-instruct-4bit`, `gemma3-1b-4bit` is `gemma-3-1b-it-4bit`) or genuinely gone, and mapping them needs per-name judgment: `jamba-v0.1-4bit` has no successor in the store and the nearest name is a different model. Left as a follow-up rather than guessed, since a wrong mapping points a parity test at the wrong checkpoint, which is worse than skipping.
inureyes added a commit that referenced this pull request Sep 9, 2026
…1718

The comment introduced with this change called the float32 promotion "a second, local copy of the same defense" as #1718, which reads as "that fix covers this site, so removing this is safe by inference". It does not, and the PR body says the opposite.

Checked in the bridge: #1718 routes `max_all`, `min_all`, `sum_all`, `mean_all` and their axis forms through `reduce_in_f32` (`mlx_cxx_bridge.cpp:840`), while `fast_rms_norm` (line 3675) calls `mlx::core::fast::rms_norm` directly and never reaches those wrappers. Two code paths that do not meet.

What actually covers the removal is `tests/mamba2_hybrid_finite.rs`, 250 forwards per checkpoint, 3 runs of 3 on both hosts with the promotion removed, plus a 3000-token needle recall on M5 Max. The comment now says that, and says the measurement has to be repeated rather than inferred if MLX changes the kernel. Comment only.
inureyes added a commit that referenced this pull request Sep 9, 2026
…de (#1724)

Three families carry the same local float32 promotion against the M5 Max half-precision `x^2` NaN. #1723 removes it from `granitemoehybrid.rs` for 1.079x. The other two were measured as well, and both answers are "leave it", so nothing in the code moved and the finding had nowhere to live. Written where the next person will look, which is the struct itself, not an issue.

`nemotron_h.rs`: measured on M1 Ultra with nvidia-nemotron-3-nano-30b-a3b-4bit, five runs per arm, 96.51 against 97.31 tok/s. Ranges are disjoint but the gain is 1.008x, and 0.8% does not buy giving up a NaN defense whose coverage after #1718 is not established, since that fix changed the bridge's reduction helpers and not MLX's fused `fast_rms_norm`.

`falcon_h1.rs`: there is nothing to measure rather than nothing measured. `mamba_rms_norm` defaults to false and the only Falcon-H1 checkpoint in the store spells it `false`, so the gated-norm `forward` is never reached on any checkpoint available to test with.
inureyes added a commit that referenced this pull request Sep 9, 2026
* perf(granitemoehybrid): run the gated RMSNorm in the input dtype

`GraniteMoeHybridRMSNormGated::forward` promoted the whole computation to float32 and cast back at the end. The promotion escapes through the mixer: the residual stream widens to f32 and each later matmul promotes its own weight to match, on every decode step.

Measured on M1 Ultra with `mlxcel-bench-decode`, before and after alternating, five runs each at `-n 200 --ignore-eos`, Time Machine confirmed stopped. granite-4.0-h-350m-4bit 272.29 to 293.88 tok/s (1.079x, ranges 265.46-273.10 and 286.78-294.35, disjoint). granite-4.0-h-tiny-4bit 107.43 to 112.17 tok/s (1.044x, ranges 105.51-107.78 and 109.95-112.74, disjoint).

Output changes, and the teacher-forced logit trace says where. At decode width (1 position, 512 tokens of context, 400 positions) top-1 disagreement is 19/400 = 4.75% overall and 0/156 on decided positions; at prefill width 256 over 1024 positions it is 51/1024 = 4.98% overall and 0/386 decided. Every disagreement sits in the undecided band, worst rank 3, perplexity +0.145% and +0.410%. The control, the baseline binary against itself at the same width, is 0/400.

The float32 promotion was a local defense against a NaN in a half-precision `x^2` sum on M5 Max. #1718 found that fault in the bridge's reduction helpers and fixed it there, but it did not touch MLX's fused `fast_rms_norm`, so this removal is not covered by that fix and rests on the guard instead. `tests/mamba2_hybrid_finite.rs` passes 3 runs of 3 with the arm applied, on M1 Ultra, where the fault does not reproduce. The M5 Max half is required before merge.

* docs(granitemoehybrid): say the removal rests on measurement, not on #1718

The comment introduced with this change called the float32 promotion "a second, local copy of the same defense" as #1718, which reads as "that fix covers this site, so removing this is safe by inference". It does not, and the PR body says the opposite.

Checked in the bridge: #1718 routes `max_all`, `min_all`, `sum_all`, `mean_all` and their axis forms through `reduce_in_f32` (`mlx_cxx_bridge.cpp:840`), while `fast_rms_norm` (line 3675) calls `mlx::core::fast::rms_norm` directly and never reaches those wrappers. Two code paths that do not meet.

What actually covers the removal is `tests/mamba2_hybrid_finite.rs`, 250 forwards per checkpoint, 3 runs of 3 on both hosts with the promotion removed, plus a 3000-token needle recall on M5 Max. The comment now says that, and says the measurement has to be repeated rather than inferred if MLX changes the kernel. Comment only.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix: Resolve flaky text_only_forward_produces_finite_logits in granite4_vision and hunyuan_vl parity tests

1 participant