Fix compiled kernel correctness for negative-strided inputs - #3720
Conversation
db1bdb0 to
ac40f52
Compare
| if (non_scalar_inputs > 1 && !all_row_contig && !all_col_contig) { | ||
| contiguous = false; | ||
| } else if (non_scalar_inputs == 1 && !all_contig) { | ||
| } else if (non_scalar_inputs == 1 && !(all_row_contig || all_col_contig)) { |
There was a problem hiding this comment.
Why the change here? Note that all_contig != (all_row_contig || all_col_contig)
There was a problem hiding this comment.
A negative-strided array (e.g. x[::-1]) has flags().contiguous == true (densely packed, no gaps) but flags().row_contiguous == false. The old check let these through to the contiguous compiled kernel, which iterates linearly from the data pointer — but for a reversed view the data pointer is at the last element, so the linear iteration reads the wrong values. This aligns the single-input path with the multi-input check on line 103 (which already uses row_contiguous / col_contiguous).
There was a problem hiding this comment.
Yeah I know that but for a single contiguous input with positive strides we don't need to use the strides at all which is why that check was like that.
Basically now doing gelu on a transposed array will be slower than before for no reason.
angeloskath
left a comment
There was a problem hiding this comment.
Thanks for the fix!
Did a little change, feel free to take a look, but basically gelu(x.transpose(0, 2, 1)) will now route to the appropriate contiguous kernel.
2069dbf to
7bb86ee
Compare
The compiled (fused) kernel path produced wrong results when inputs had negative strides (e.g. x[::-1]). Two issues: 1. compiled_check_contiguity used the broad contiguous flag for single inputs, which is true for negative-strided arrays (no data gaps). Changed to require row_contiguous or col_contiguous, matching the multi-input path. 2. Metal/CUDA strided compiled kernels used unsigned index arithmetic (elem_to_loc_1<uint>), wrapping negative strides. Force int64_t indices when any input has negative strides. Also generate the _large (int64_t) strided kernel variant for ndim=1. The CPU compiled path uses signed pointer arithmetic and only needed the contiguity check fix. Fixes ml-explore#3716. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
The strides vector from compiled_collapse_contiguous_dims is ordered [output, input_0, input_1, ...]. Expand the comment to make this clear and explain why only input entries need the negative-stride check. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
7bb86ee to
5c23944
Compare
…-strides API Upstream b5404c9 changed compiled_collapse_contiguous_dims to return a 4-tuple whose new second element reports negative input strides, and made the helper route any negative-strided case to the non-contiguous path. Destructure the new tuple in the WebGPU Compiled::eval_gpu. No extra routing is needed here: Metal/CUDA force 64-bit "large index" mode because their small-index kernels did unsigned index math, but the WebGPU strided kernel is already sign-correct — strides bind as array<i32>, elem_to_loc accumulates a signed i32 loc, and the read site adds the (element) base offset in i32 before the final u32 cast, which is exact because offset + sum(idx_d * stride_d) >= 0 for any in-bounds element. WGSL has no i64, so large mode cannot exist on this backend. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
## Summary - Bump all three synchronized MLX pin locations to upstream `57c66cac7cb3e5b1eb350488a61f1506b40d39f8` (`Patch bump to 0.32.1`, ml-explore/mlx#3816). - Fix mlxcel's Metal compiled patch overlay for the latest `compiled_collapse_contiguous_dims` return shape by binding `negative_strides`, selecting the large-index path for negative strides, and generating the rank-1 large strided Metal kernel that upstream MLX now needs. - Note the relevant upstream range includes ml-explore/mlx#3804 (`Fix fp quantized matvec for output dim < 8`), ml-explore/mlx#3728 (`Add math mode option for custom Metal kernels`), and ml-explore/mlx#3720 (`Fix compiled kernel correctness for negative-strided inputs`). ## Test plan - [x] `cargo build --release --bin mlxcel-bench-decode` - [x] `cargo test --release -p mlxcel-core sparse_v_kernel_threshold_zero_matches_graph` - [x] `cargo test --release -p mlxcel-core delegated_fused_kernel_matches_reference_over_200_steps` - [x] `cargo test --release -p mlxcel-core delegated_steel_envelope_matches_cold_only_fused_over_200_steps` - [x] `python3 scripts/ci/check_cross_repo_refs.py` Closes #703
…-strides API Upstream b5404c9 changed compiled_collapse_contiguous_dims to return a 4-tuple whose new second element reports negative input strides, and made the helper route any negative-strided case to the non-contiguous path. Destructure the new tuple in the WebGPU Compiled::eval_gpu. No extra routing is needed here: Metal/CUDA force 64-bit "large index" mode because their small-index kernels did unsigned index math, but the WebGPU strided kernel is already sign-correct — strides bind as array<i32>, elem_to_loc accumulates a signed i32 loc, and the read site adds the (element) base offset in i32 before the final u32 cast, which is exact because offset + sum(idx_d * stride_d) >= 0 for any in-bounds element. WGSL has no i64, so large mode cannot exist on this backend. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Review of the pin bump found three things the new pin changes underneath this tree. - ml-explore/mlx#4208 moved Cholesky onto cuSOLVER and `gpu::init()` now creates its handle cache on every CUDA start. MLX links it PRIVATE, so cargo never saw it and every `--features cuda` link would fail on `cusolverDnCreate`; `link_cuda()` now names `cusolver`. - ml-explore/mlx#3742 made `array::is_available()` detach the event through `Event::check_error()`, which throws and clears a failed launch's error. The rejection sampler's deferred drain called it on slots other requests stashed, inside `fused_sample`, which is not a `Result` bridge, so a failed command buffer would terminate the process and hide the error from the request that owns it. The drain now reads status, signal and error pointer directly and drops a failed slot unread. - The Metal `compiled.cpp` overlay still emitted `elem_to_loc_1<uint>` for 1-D inputs, half of ml-explore/mlx#3720 that an earlier sync missed; it now matches upstream, so the overlay's only delta is the mixed-dtype cast. The CUDA mixed-type `FloorDivide` overload floors like upstream's float branch (ml-explore/mlx#4108), and three stale sync notes are corrected. Workspace gate 10985 passed, 0 failed; clippy and fmt clean on Metal. The CUDA link is not verifiable on this host. Refs #1769
) ## Why The fp8 round-trip bound failed on every Metal host, M1 Ultra byte-identically to M5 Max, because the pinned MLX `9a795735` predates ml-explore/mlx#4353: Metal and CPU encoded the mxfp8 E8M0 block scale as `round(log2(amax / 448))`, so about half the blocks saturated their maxima, losing up to `1 - 2^-1/2`. CUDA rounds up, which is why #1742 passed on GB10. The test was right, and `requantize_block_fp8_weights`, the only E8M0 quantize caller, was clipping vendor FP8 checkpoints on Metal. Widening the bound, as the issue proposed, was rejected. ## What changed - MLX pin `9a795735` to upstream main `81ba1c6a` (99 commits). The seven overlays whose targets upstream touched are three-way merged and keep their deltas; the other 21 are unchanged upstream. `metal/compiled.cpp` also drops a leftover `elem_to_loc_1<uint>` that undid part of ml-explore/mlx#3720, so its only delta is the mixed-dtype cast. - Adaptations to what the new pin changes under the bridge: - `gather_qmm` gained `global_scale` ahead of `sorted_indices` (ml-explore/mlx#4458), so all 13 calls pass `std::nullopt`. - CUDA now needs cuSOLVER (ml-explore/mlx#4208). `link_cuda()` names it, `docs/installation.md` lists it, and CI's link job now runs on pin and CUDA link-list changes, which is how this slipped past CI. - `array::is_available()` now throws and clears a failed launch's error (ml-explore/mlx#3742). The rejection sampler's drain now reads status, signal and error pointer instead, so a GPU fault no longer terminates the process from inside `fused_sample` or hides the error from its owner. A regression test aborts the process with the old drain and passes with the new one. - The round-trip check moves into its own test, `fp8_block_requantize_round_trip_stays_within_half_an_e4m3_step`, with the same seed and shape. It states the derivation and asserts `group_max <= 448 * scale` per block. ## Validation (M1 Ultra, macOS 27.0) - fp8: the old pin fails the new test at block 0 (the maximum 4.8046875 scales to 615 and saturates). The new pin saturates 0 of 650 blocks, with a worst error of 0.0489 of the group max against 0.2928 before. - Turbo launchers pass (max RMS 1.7263e-4 and 1.5259e-4). - Teacher-forced logit traces, old pin vs new, on five checkpoints at widths 1, 8 and 256: 0 disagreements on decided positions. Over 4,096 positions on qwen3-30b-a3b, 1 of 1,720 decided positions differs, at the reference's rank 2, with perplexity -0.20%. - Branches the short runs missed: - head-dim-512 decode past 1,024 keys: identical text. - GQA-8 decode past 8,192 keys: identical text, decode 55.6 to 60.8 tok/s. - head-dim-72 vision towers: image prefill 239 to 279 tok/s. Descriptions diverge into equally faithful text; decided answers are unchanged. - Short-context throughput is within 0.6% on three checkpoints. - Workspace gate: 10986 passed, 0 failed, 359 ignored. Clippy, fmt, both pin parsers and the cross-repo reference check are clean. - CUDA was compiled and linked in CI but not run: `OpenXLA feature link` linked a `--features cuda,xla-iree` release binary on GB10 at the new pin, which covers the overlays and the cuSOLVER link. The green `CUDA sm_70 compile` check is a skip, because CUDA 13.0 cannot target sm_70. The per-checkpoint numbers, method and derivations are in `TECHNICAL_REPORTS/1772-mlx-pin-mxfp8-round-up-20260911.en.md`. ## Not validated, or reported and not fixed - No CUDA execution. No M5 Max run: the generation-17 NAX paths changed upstream in this range are unmeasured. - `array_evaluated_bytes`, the server's lookahead read, is another non-`Result` bridge function that now throws on a failed launch. It needs routing through the scheduler's step-failure path. - The mixed-dtype cast in `metal/compiled.cpp` would cast a comparison's inputs to `bool`. This is latent, since no compiled function contains a comparison. Closes #1769
* ci: pin clang-tidy to LLVM 22 Homebrew's `llvm` formula moved to 23.1.0, whose clang-tidy adds readability-trailing-comma, readability-redundant-nested-if, readability-redundant-lambda-parameter-list, misc-explicit-constructor and bugprone-signed-bitwise to the enabled set. That turned Lint C++ red on PRs that touch no C++ (e.g. #233) with ~21 findings in untouched files, while `main` stayed green only because it had an older cached Homebrew archive. Pin the tool to `llvm@22`, the version `main` last passed with, so lint is reproducible and new checks are adopted in a deliberate commit. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * Bump mlx to 0.32.0 * Patches updated accordingly. * test(eigh): trigger the worker-thread failure directly, not via LOBPCG The LOBPCG trigger was itself a symptom of #232. Its Rayleigh-Ritz step reverses the basis (`V[:, ::-1]` in jax/experimental/sparse/linalg.py), a reversed view feeding a fused kernel read as zeros past the first element, and the next iteration's syevd got a degenerate projection and threw. With that MLX bug fixed, LOBPCG converges and the test stops exercising the worker-thread path patch 10 exists for -- which its own NOABORT:COMPLETED assertion is designed to catch. Feed eigh an all-NaN matrix instead: syevd returns info != 0 regardless of MLX version, so the trigger no longer depends on a miscompile. Verified the failure still surfaces as a catchable Python exception carrying "Eigenvalue decomposition failed". Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * test(shape): cover a reversed view feeding a broadcast operand A negative-stride view meeting a broadcast operand in the same fused MLX kernel read as zeros past the first element (#232): `jnp.flip(y) * 2.0` returned `[16, 0, 0, ...]` on MPS. Bare `jnp.flip` -- the only reversal the suite exercised -- stayed correct throughout, so nothing caught it. Fails against the previous MLX pin (2 failed, "Values are not close") and passes on v0.32.0, which carries the upstream fix (ml-explore/mlx#3720). Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com> Co-authored-by: MasterSkepticista <karanshah1698@gmail.com>
tests/link/mlx/conftest.py pins the default MLX device to CPU for the whole tests/link/mlx/ directory. The negative-strided mx.compile miscompile of ml-explore/mlx#3720 only affects the GPU stream, so the two tests marked xfail for MLX <0.32.0 now pass on the pinned CPU device and the strict markers turn into XPASS failures. Remove both markers. MLX supports float64 on CPU, so the check no longer runs in float32 either; drop the loosened rtol and the ScanCompatibilityTests.check_higher_order_derivative parameter that existed only to carry it.
Fixes #3716.
This is an alternative to #3717 that fixes the compiled-kernel path in place instead of materializing negative-strided inputs before the generated kernel runs.
mx.compileproduces incorrect results for negative-strided inputs such asx[::-1], while eager execution is correct.This patch fixes two compiled-kernel path issues:
_largeMetal variant is now always generated (matching ndim > 1 behavior) since negative strides force int64_t indexing even for 1D arrays.Added regression coverage for 1D reverse, slice update, 2D reverse, mixed positive/negative strides, and a 4D negative-stride case.
Tests:
python -m pytest python/tests/test_compile.py -q # 59 passed