Skip to content

fix(rocm): guard the four remaining custom-kernel launchers - #2018

Merged
inureyes merged 2 commits into
mainfrom
fix/rocm-remaining-kernel-launcher-guards
Sep 28, 2026
Merged

inureyes merged 2 commits into
mainfrom
fix/rocm-remaining-kernel-launcher-guards

Conversation

@inureyes

Copy link
Copy Markdown
Member

Refs #1885, #1814, #1801.

Issue #1885 fixed this defect at the two sampler launchers. The same shape survived at four more: a two-valued dispatch whose false arm still meant "Metal" rather than "no port for this backend", so a direct call on ROCm took the Metal arm and fast::metal_kernel threw across a noexcept cxx extern, ending the process.

Classifying all ten dispatch sites mechanically found them: five were already guarded by #1803 and #1885, four were not, and the MoE launcher is reached through two bridge entry points. The guards reuse the shape paged_attention_decode established, placed before the port is selected so the message names the backend rather than the port that happened to be tried, and five bridge declarations become Result in the same change, because a guard without that conversion only changes which abort message prints.

Each guard was checked against its production gate before being added, since a guard above a reachable path would turn a working graph fallback into a throw. fused_add_rms_norm_available() and fused_rope_qk_append_available() are custom_kernels_available(); ssm_kernel_available() is cu::is_available(), which is the no-CUDA stub on a ROCm build; fused_moe_enabled() folds in the port predicate. All four answer false on ROCm, so production takes its graph path and never reaches a guard.

The ten callers split three ways, and the split is deliberate. The MoE sites use .ok() because the function wrapping them returns Option and "the fast path does not apply" is already its contract, so a direct call on a portless backend is legitimate rather than a defect. The SSM, layers.rs and parity-test sites use expect because each already gates on the support predicate, so an error there means the gate and the launcher disagree, which should fail loudly instead of being folded into a fallback.

Two comments this change made false are corrected rather than left to mislead. fused_moe_enabled() justified its port term with "its bridge function does not return Result", which was true under #1803 and is not now; the term stays, for the cheaper reason that deciding early beats building the call to have it refused. fused_add_rms_norm_eligible() claimed the launcher's exceptions are not recoverable, which now holds for contract violations but not for a missing port.

Validated on gfx1151: the full ROCm gate is unchanged at one failing target, the nvfp4 abort of #1806, with that abort the only terminate called in the run. A probe confirmed the refusal by running it, returning [fused_add_rms_norm] no custom kernel port for this GPU backend as a typed error rather than aborting. CUDA and Metal are unreachable from every new path, since custom_kernels_available() is true there.

The ten dispatch sites

Classified mechanically rather than by reading, so none was missed.

Guarded already Unguarded, fixed here
gumbel_max_sample (#1885) fused_add_rms_norm
rejection_sample (#1885) fused_rope_qk_append
paged_attention_decode (#1803) ssm_update_kernel
paged_attention_decode_v2_partial run_fused_moe_two_kernel
paged_attention_merge_states

run_fused_moe_two_kernel is reached through two bridge entry points, so five declarations become Result, not four.

Evidence

The refusal, run rather than inferred:

custom_kernels_available = false
fused_add_rms_norm: Err -> [fused_add_rms_norm] no custom kernel port for this GPU backend; mlxcel's callers take the graph fallback instead

Full ROCm gate, unchanged from before this branch:

error: 1 target failed:
    `-p mlxcel-core --lib`
aborts: 1

That one is the nvfp4 abort of #1806. cargo fmt is clean, and cargo clippy -p mlxcel --lib --tests -- -D warnings, the command CI runs, is clean.

Worth noting how two of the ten callers were found. cargo check passed on them, because ignoring a Result<()> is the unused_must_use lint rather than a type error, so only the -D warnings clippy surfaced them. Both are decode hot-path callers in layers.rs.

Next

Every guard here is a placeholder that says "no port". The ports themselves are #1814, and the path is already open: fast::hip_kernel exists in the ROCm overlay and #1862 proved the three-arm switch with BitNet.

Issue #1885 fixed this defect at the two sampler launchers. The same shape survived at four more: a two-valued dispatch whose false arm still meant "Metal" rather than "no port for this backend", so a direct call on ROCm took the Metal arm and `fast::metal_kernel` threw across a `noexcept` cxx extern, ending the process.

Classifying all ten dispatch sites mechanically found them: five were already guarded by #1803 and #1885, four were not, and the MoE launcher is reached through two bridge entry points. The guards reuse the shape `paged_attention_decode` established, placed before the port is selected so the message names the backend rather than the port that happened to be tried, and five bridge declarations become `Result` in the same change, because a guard without that conversion only changes which abort message prints.

Each guard was checked against its production gate before being added, since a guard above a reachable path would turn a working graph fallback into a throw. `fused_add_rms_norm_available()` and `fused_rope_qk_append_available()` are `custom_kernels_available()`; `ssm_kernel_available()` is `cu::is_available()`, which is the no-CUDA stub on a ROCm build; `fused_moe_enabled()` folds in the port predicate. All four answer false on ROCm, so production takes its graph path and never reaches a guard.

The ten callers split three ways, and the split is deliberate. The MoE sites use `.ok()` because the function wrapping them returns `Option` and "the fast path does not apply" is already its contract, so a direct call on a portless backend is legitimate rather than a defect. The SSM, `layers.rs` and parity-test sites use `expect` because each already gates on the support predicate, so an error there means the gate and the launcher disagree, which should fail loudly instead of being folded into a fallback.

Two comments this change made false are corrected rather than left to mislead. `fused_moe_enabled()` justified its port term with "its bridge function does not return `Result`", which was true under #1803 and is not now; the term stays, for the cheaper reason that deciding early beats building the call to have it refused. `fused_add_rms_norm_eligible()` claimed the launcher's exceptions are not recoverable, which now holds for contract violations but not for a missing port.

Validated on gfx1151: the full ROCm gate is unchanged at one failing target, the nvfp4 abort of #1806, with that abort the only `terminate called` in the run. A probe confirmed the refusal by running it, returning `[fused_add_rms_norm] no custom kernel port for this GPU backend` as a typed error rather than aborting. CUDA and Metal are unreachable from every new path, since `custom_kernels_available()` is true there.

Refs #1885, #1814, #1801.
@inureyes inureyes added type:bug Bug fixes, error corrections, or issue resolutions status:review Under review priority:medium Medium priority labels Sep 28, 2026
@inureyes
inureyes merged commit e2a2ebe into main Sep 28, 2026
19 of 24 checks passed
@inureyes
inureyes deleted the fix/rocm-remaining-kernel-launcher-guards branch September 28, 2026 14:14
inureyes added a commit that referenced this pull request Sep 28, 2026
* refactor(core): choose custom-kernel ports through one helper

Choosing a fused kernel's port was hand-written at every launcher, as `use_cuda ? cuda_port : metal_port`. That reads "not CUDA" as "Metal", which was true while Metal and CUDA were the only backends and became a defect the moment a third existed: on ROCm the false arm was taken, `fast::metal_kernel` threw, and because the bridge declarations were not `Result` the throw crossed a `noexcept` cxx extern into `std::terminate`.

The history is the argument for this change. #1803 replaced the underlying `!metal::is_available()` test with a named backend kind and left nine sites still shaped that way. #1885 and #2018 then guarded those nine by hand, one launcher at a time, months apart, each after it had already aborted a gate run. The hand-written guards drifted: two cited support predicates that do not exist, one hardcoded an entry-point name that was wrong for one of its two callers, and the wording differed per site. A convention that has to be re-applied at every new launcher gets missed at some of them, and only the backend that lacks the port ever finds out.

So the choice now goes through one helper. A kernel declares its ports once as a `KernelPorts` table, `select_kernel_port` owns the refusal and its message, and `has_kernel_port` answers the support predicate from the same table. That last part removes a real inconsistency rather than a hypothetical one: `ssm_kernel_available()` answered with `cu::is_available()` while its launcher guarded on `custom_kernels_available()`, two predicates for one question, free to disagree. A caller that gates on the table may now `expect` the launch by construction.

A null table entry is "no port for that backend", which makes absence a value rather than something a dispatch falls into. It also gives a Metal-only kernel the same shape as a ported one, and gives #1814 a named slot to fill.

Three defenses replace the convention, and each was confirmed by negative control rather than assumed. `scripts/ci/check_kernel_port_dispatch.py` rejects a reintroduced backend comparison and a hand-rolled refusal, both verified by putting them back. A `static_assert` on `kGpuKernelBackendCount` fails the build when a backend is added, with a message naming what to extend, verified by adding a `Vulkan` enumerator; the assert exists rather than relying on `-Wswitch` because this repository compiles with warnings on but not as errors, and a new one would be lost among the MLX warnings. The checker runs in `verify` and `verify-rocm` and as an unconditional hosted CI job, for the reason the neighbouring `kernel dtype keys` job gives: no toolchain, seconds to run, and the defect it guards is silent.

All nine two-way launchers are converted, not a subset, because a checker that has to exempt half of them cannot be strict and an exemption list tends to become permanent. Two Metal-only SDPA launchers are exempt and listed with their reason: neither has any refusal today, and `sparse_v_available()` gates on an env threshold and a KV cache mode with no backend term, so adding a guard there is a behavior change needing its own reachability analysis rather than a mechanical edit.

Fixed in passing, because the standard form makes it impossible: `run_fused_moe_two_kernel` now takes its caller's entry-point name, so a refusal reached through `fused_moe_geglu_kernel` no longer reports `fused_moe_expert_kernel`.

Validated on gfx1151: the full ROCm gate is unchanged at one failing target, the nvfp4 abort of #1806, the only `terminate called` in the run. `cargo fmt`, the CI clippy command, and `actionlint` are clean.

Refs #1801, #1803, #1885, #2018, #1814.

* docs: add technical report for PR #2026
inureyes added a commit that referenced this pull request Sep 29, 2026
* refactor: finish the kernel port standardization

PR #2026 routed nine two-way launchers through `select_kernel_port`. This closes the three gaps it left and corrects one claim it recorded. Full reasoning in the updated `TECHNICAL_REPORTS/2026-*` pair.

Checker rule 3, a launcher must not reach a kernel holder directly. Rules 1 and 2 only describe how a multi-port launcher goes wrong; a single-port one goes wrong by calling `get_x_kernel().get()` with no resolution, which neither sees. Found by reverting a Metal-only launcher by hand and watching the check still pass. It then found five launchers every earlier survey had missed, so the exemption list is now empty.

Checker rule 4, a Rust gate must not be spelled `metal_is_available() || cuda_is_available()`. The condition is not wrong today, and that is the problem: it names the two backends that happen to have ports, so on a third it reads as a missing term, and "add `rocm_is_available()`" is the natural conclusion and the wrong one. Both kernels' `.rocm` entries are still null, so widening the gate would run the caller past the port table into the launcher's refusal. Two questions shared one name. "Has this kernel's port" is now the kernel's own `*_available()`, reading `has_kernel_port` on the table the dispatch reads; "is there a GPU at all" is now `gpu_backend_available()`. Rule 4's two exemptions are missing predicates, not deferred convention, each tracked to #1814.

Not cosmetic: `rms_norm_small_axis_tests.rs` tests MLX's own `fast::rms_norm` dispatch config, not an mlxcel port, so the narrow gate had been skipping it on ROCm for no reason. Under the correct predicate both sweeps run on gfx1151 and pass, which is new coverage.

The SDPA correction: #2026's report called the two SDPA launchers' reachability off Apple unresolved and treated a latent abort as likely. Wrong. Both are reached only through `sparse_v::kernel_enabled()` (`sparse_v.rs:151`), false unconditionally off macOS, verified by forcing the sparse-V branch with `MLXCEL_TURBO4_ASYM_DEQUANT_SDPA=0` and generating normally. They join the table as consistency, not as a bug fix.

Three `Result<()>` values dropped from launchers #2018 converted are real swallowed refusals. The `fused_xielu` gate in `apertus.rs` is not one: `mlx_cxx_kernels.cpp:463` early-returns an elementwise fallback, so it falls back rather than refusing, and an earlier change here was reverted. The `nemotron_h.rs` gate is defensive, not a fixed bug: that path was traced, not executed, for want of a checkpoint on this host.

Why five such sites survived four rounds: an ignored `Result<()>` is only `unused_must_use`, so `cargo check` stays green, and `cargo clippy -p mlxcel --lib --tests`, which PR-time CI runs, builds neither mlxcel-core nor examples. `make verify` and `make verify-rocm` lint `--workspace --all-targets` and found all of them. `kernel_port.h` now says so beside the declaration requirement.

Validation on gfx1151: the ROCm gate is unchanged at one failing target, `-p mlxcel-core --lib`, the nvfp4 abort of #1806, and that abort is the run's only `terminate called`. It kills the test binary before anything alphabetically later runs, so the newly enabled sweeps were run directly rather than inferred. Workspace clippy and `fmt --check` clean. All four checker defenses confirmed by negative control, rule 4 in three spellings. `verify-rocm-smoke` could not run: its fixture lived under `/tmp`, cleared on reboot, and no local checkpoint remains.

Refs #1801, #1803, #1814, #1885, #2018, #2026

* docs: add technical report for PR #2029
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

priority:medium Medium priority status:review Under review type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant