Skip to content

Fix compiled kernel correctness for negative-strided inputs - #3720

Merged
zcbenz merged 3 commits into
ml-explore:mainfrom
lyonsno:cc/fix-compile-negative-stride
Jul 7, 2026
Merged

zcbenz merged 3 commits into
ml-explore:mainfrom
lyonsno:cc/fix-compile-negative-stride

Conversation

@lyonsno

@lyonsno lyonsno commented Jun 19, 2026 •

Copy link
Copy Markdown
Contributor

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.compile produces incorrect results for negative-strided inputs such as x[::-1], while eager execution is correct.

This patch fixes two compiled-kernel path issues:

  • single-input contiguity detection now requires row- or column-contiguous layout instead of the broader contiguous flag, so negative-strided views route to the strided compiled path;
  • Metal/CUDA compiled strided kernels use signed 64-bit indexing when an input has negative strides. The ndim=1 _large Metal 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

Comment thread mlx/backend/cuda/compiled.cpp Outdated
Comment thread mlx/backend/common/compiled.cpp Outdated

@zcbenz zcbenz left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks!

Comment thread mlx/backend/common/compiled.cpp Outdated
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)) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why the change here? Note that all_contig != (all_row_contig || all_col_contig)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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).

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 angeloskath left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@angeloskath
angeloskath force-pushed the cc/fix-compile-negative-stride branch from 2069dbf to 7bb86ee Compare July 6, 2026 18:18
lyonsno and others added 3 commits July 6, 2026 15:47
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>
@angeloskath
angeloskath force-pushed the cc/fix-compile-negative-stride branch from 7bb86ee to 5c23944 Compare July 6, 2026 22:47
@zcbenz
zcbenz merged commit b5404c9 into ml-explore:main Jul 7, 2026
28 checks passed
@BrewTestBot BrewTestBot mentioned this pull request Jul 7, 2026
1 task done
Brooooooklyn added a commit to mlx-node/mlx that referenced this pull request Jul 8, 2026
…-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>
inureyes added a commit to lablup/mlxcel that referenced this pull request Jul 9, 2026
## 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
Brooooooklyn added a commit to mlx-node/mlx that referenced this pull request Jul 22, 2026
…-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>
inureyes added a commit to lablup/mlxcel that referenced this pull request Sep 10, 2026
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
inureyes added a commit to lablup/mlxcel that referenced this pull request Sep 10, 2026
)

## 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
tillahoffmann added a commit to tillahoffmann/jax-mps that referenced this pull request Sep 12, 2026
* 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>
cetagostini added a commit to cetagostini/pytensor that referenced this pull request Sep 15, 2026
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.
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.

mx.compile: assigning an elementwise expression to a negative-strided slice writes only one element

3 participants