Skip to content

update(rocm): port the Gumbel-max and rejection samplers to HIP - #2100

Merged
inureyes merged 7 commits into
mainfrom
update/issue-2064-hip-samplers
Oct 4, 2026
Merged

inureyes merged 7 commits into
mainfrom
update/issue-2064-hip-samplers

Conversation

@inureyes

@inureyes inureyes commented Oct 4, 2026 •

Copy link
Copy Markdown
Member

Summary

Both fused samplers took the MLX graph on ROCm: gumbel_ports() and rejection_ports() had no .rocm entry, and rejection_sample_supported() answered with custom_kernels_available() (Metal or CUDA by definition), so a filled slot would have stayed unreachable. This fills both slots and makes the rejection predicate read the port table, so fused_sample, fused_sample_probs and the speculative paths reach the HIP kernels with no caller change.

  • sampling_gumbel_hip.h, sampling_rejection_hip.h: the CUDA bodies with the same inputs, outputs, grid, template arguments (dtype keys included) and Philox-4x32-10 counter and key layout. Spelling only: explicit (float) reads (hip_bfloat16 converts only explicitly) and -__builtin_huge_valf() for -INFINITY. The hip_kernel holders and launches stay in sampling.cpp / sampling_rejection.cpp, so EXPECTED_IN_SCOPE and its count are unchanged.
  • No wave32 #error: neither kernel has a lane-level operation (every reduction and the scan are shared memory with a barrier per step), so both are correct on wave64, and the guard idiom is inert on ROCm 10's clang anyway (perf(rocm): port ssm_update_kernel to HIP and fix ssm_kernel_available #2067). A guard here would catch nothing.
  • rejection_sample_supported() returns has_kernel_port(rejection_ports()); still true on Metal and CUDA. The two kill-switch integration tests chose their expected decline with custom_kernels_available(), which stays Metal-or-CUDA, so they now ask new env-independent bridge predicates (sampling_{gumbel,rejection}_backend_supported()); the stale "ROCm has no ports" comments and the MLXCEL_DEBUG_KERNEL_BACKEND suffix are reworded.
  • Overlay fix (LOCAL_FIXES item 30, upstreaming candidate for chore(rocm): mlxcelverse ROCm fork sync script, MLX pin-bump procedure, and upstreaming local fixes #1813): the vendored fast::hip_kernel declared <input>_shape / _strides as pointers while the launch passes them by value, so the Gumbel port's logits_shape[1] faulted the queue. They are now by-value structs, as upstream CUDA's __grid_constant__ Shape is, and the launch appends them only for ndim > 0, matching the declaration. No earlier HIP kernel read either.
  • bench_decode.sh --temperature/--top-p (plain decimals only, since they tag the output filename), so sampled decode can be measured with the standard harness.

Results on gfx1151

Llama-3.1-8B-Instruct-4bit, pp512/tg128, three runs per arm, before (57d8ed29) and after interleaved, each through rocm_gpu_guard.sh (the #2098 copy from the #2065 worktree, not committed here; --idle-secs 75, 112 to 116 s of wall-clock idle per run, chosen to stop tying with a parallel unit's 90):

Config Before median After median
--temperature 0.8 (Gumbel) 37.91 38.26 inside the noise (1.7% / 4.4% spread)
--temperature 0.8 --top-p 0.95 (rejection) 33.98 37.56 1.11x; per pair 1.04x / 1.11x / 1.13x, the before arm drifted down

--top-k 40 --top-p 0.95 is not routed to the kernel at vocab 128256 (joint cap 32768), so it was not measured. Details: docs/benchmark_results/rocm-samplers-gfx1151-2026-10-05.md.

Verification

On gfx1151 (Radeon 8060S, ROCm 10.0.0):

  • make verify-rocm at 6e336ec8: OK; 11861 passed, 0 failed, 378 ignored across 146 test binaries; ROCm smoke OK. The one later commit (74624ea7) changes docs, two header comments and bench_decode.sh input validation; the fast gates, fmt and dead_doc_pointers pass on it.
  • The gate before it, at f7bec56d, failed only temperature_one_support_unchanged. That test skipped its two routed cases where the rejection kernel had no port, so they first ran on ROCm here. Those rows are the graph's softmax over the kernel's support (fused_sample_probs never launches the kernel); gfx1151 lands 2 ulp from the Metal capture at one token with the support unchanged, so the bound goes from 1 to 2 ulp (the stock case allows 4).
  • cargo test --release --features rocm -p mlxcel-core --lib sampling_ -- --test-threads=1: 53 passed, plus --test sampling_gumbel_kill_switch --test sampling_rejection_kill_switch; sampling_gumbel_tests and sampling_rejection_tests now run instead of returning early.
  • sampling_fixed_key_tests: kernel draws recomputed on the host from the consumed key (Gumbel also through the MLX graph, f32/f16/bf16; rejection under min-p), rejection draws under top-p and top-k (multi-round rows included) inside the host-computed support, and two-sample chi-square tests against fused_sample_categorical. Without the overlay fix the first test faults the queue; with the Philox counter or word changed in either HIP source the fixed-key tests fail.
  • Speculative decode (no MTP checkpoint here): Qwen3-30B-A3B with a Qwen3-0.6B drafter at --temp 0.8 --top-p 0.95, both acceptance rules, before and after: the log names the rejection kernel after, output is coherent, acceptance sits near its closed form.
  • cargo clippy -p mlxcel-core --features rocm --lib --tests -- -D warnings, make verify-versions verify-kernel-dtype-keys verify-kernel-port-dispatch verify-llama-compat verify-rocm-overlay verify-fmt: pass. The dtype-key checker's own test moves both launches in sampling.cpp now.

Not verified: Metal and CUDA (not available on this host). They are touched by the rejection predicate (has_kernel_port, true on both as before), the shared bridge functions (random_bits, sampling_{gumbel,rejection}_backend_supported), the kill-switch tests' predicate, the routed snapshot bound, and sampling_fixed_key_tests, which is written to run there too. No kernel reads <input>_strides, so that half of the overlay fix is untested beyond the struct compiling in every kernel header.

Closes #2064

@inureyes inureyes added status:review Under review type:performance Performance improvements priority:low Low priority area:core mlxcel-core: MLX FFI, primitives, KV cache, layers platform:linux Linux (CUDA / packaging) specific labels Oct 4, 2026
inureyes added a commit that referenced this pull request Oct 4, 2026
From review of #2100: the results page now says what the commits after the measured binary change and quotes `compare_bench_csv.py` (1.024x and 1.130x on the third runs) next to the medians and the per-pair ratios (1.04x, 1.11x, 1.13x); `docs/installation.md` notes the before arm's drift; the upstream README row for LOCAL_FIXES item 30 describes the whole change. `bench_decode.sh` rejects a `--temperature` or `--top-p` value that is not a plain decimal, since both go into the output filename. Two header comments are reflowed.

Refs #2064
@inureyes inureyes added status:done Completed and removed status:review Under review labels Oct 4, 2026
inureyes added a commit that referenced this pull request Oct 4, 2026
Bilingual (en/ko) report on the HIP ports of the Gumbel-max and rejection samplers: the line-for-line ports, the rejection predicate reading the port table, the fast::hip_kernel shape/strides overlay fix (LOCAL_FIXES item 30), the fixed-key, support and chi-square evidence, the gfx1151 measurements, and the 2 ulp snapshot bound.
The overlay's `fast::hip_kernel` declared an input's `<name>_shape` and `<name>_strides` as pointers while `CustomKernel::eval_gpu` passes them by value through `KernelArgs::append_ndim` (padded to `JIT_MAX_NDIM` entries). A kernel that read `<name>_shape` therefore dereferenced the shape values as an address and faulted the queue on its first launch; the HIP Gumbel-max sampler for #2064 hit it at `0x80000000000` on gfx1151. No HIP kernel in mlxcel used either argument before.

The generated header now defines `KernelShape` and `KernelStrides` (eight entries, tied to `JIT_MAX_NDIM` by a `static_assert`) and declares both parameters with them by value, as upstream CUDA does with `__grid_constant__ Shape` / `Strides`; `elem_to_loc` gains an overload for them. Recorded as LOCAL_FIXES item 30, an upstreaming candidate for #1813 (not packaged yet).

Refs #2064
Both fused samplers took the MLX graph fallback on ROCm because their port tables had no `.rocm` entry, and `rejection_sample_supported()` answered with `custom_kernels_available()`, which is Metal-or-CUDA by definition, so a filled slot would have stayed unreachable.

- `GUMBEL_MAX_SAMPLE_HIP_SOURCE` and `REJECTION_SAMPLE_HIP_SOURCE` (new headers, data only) are line-for-line ports of the CUDA bodies with the same inputs, outputs, grid, template arguments, Philox-4x32-10 counter and key layout. The only changes are spelling: explicit `(float)` reads, since `hip_bfloat16` converts only explicitly, and `-__builtin_huge_valf()` for `-INFINITY`. The `hip_kernel` holders and launches stay in `sampling.cpp` and `sampling_rejection.cpp`, so the dtype-key pin and its count are unchanged.
- No wave32 `#error`: neither kernel has a lane-level operation (every reduction and the scan go through shared memory with a barrier per step), so both are correct on any wavefront size.
- `rejection_sample_supported()` reads `has_kernel_port(rejection_ports())`, as `gumbel_max_sample_supported()` already did; true on Metal and CUDA as before.
- `random_bits` bridge function, and `sampling_fixed_key_tests`: each kernel's draw recomputed on the host from the key it consumed (Gumbel-max also through the MLX graph, in f32, f16 and bf16; rejection under min-p, where round 0 decides), plus two-sample chi-square tests against `fused_sample_categorical`. The fixed-key tests fail with the Philox counter or word changed in either HIP source.
- The dtype-key checker's fixture mutations now move both launches in `sampling.cpp`.

Refs #2064
`bench_decode.sh` gains `--temperature` and `--top-p`, forwarded to `mlxcel-bench-decode`, so a sampled decode can be measured with the standard harness; the CSV schema is unchanged and the auto-generated filename is tagged instead.

Llama-3.1-8B-4bit on gfx1151, three interleaved guarded runs per arm: `--temperature 0.8 --top-p 0.95` (rejection kernel) went from a median of 33.98 to 37.56 tok/s with no overlap between arms; `--temperature 0.8` alone (Gumbel-max) moved 0.9%, inside the run-to-run spread. The issue's `--top-k 40 --top-p 0.95` command is not routed to the kernel at vocab 128256 and was not measured. The results page also records the fixed-key and two-sample tests and a speculative decode check with a Qwen3 drafter; `docs/installation.md` gains a sampled-decode row.

Refs #2064
With the HIP ports, `rejection_sample_supported()` is true on ROCm while `custom_kernels_available()` stays Metal-or-CUDA, so `tests/sampling_rejection_kill_switch.rs` expected the "no rejection-sampling kernel port" decline where `fused_sample` now reports the kill switch. Both kill-switch tests now ask new env-independent bridge predicates, `sampling_gumbel_backend_supported()` and `sampling_rejection_backend_supported()`, which read the kernels' own support predicates.

Also from review:

- The overlay's custom-kernel launch appends an input's shape, strides and ndim only for ndim > 0, the condition under which `build_kernel` declares them (LOCAL_FIXES item 30).
- `sampling_fixed_key_tests` checks that rejection draws under top-p and top-k (multi-round rows included) converge and stay inside the support computed from the kernel's own probabilities, and asserts ids are in range when building histograms.
- `gpu_backend.h`, the bridge comment on `custom_kernels_available()` and the `MLXCEL_DEBUG_KERNEL_BACKEND` line no longer say ROCm has no ports; the probe-parser fixture follows the new line.
- Corrected the Gumbel HIP header note (the CUDA text already casts the logits read) and the results page wording.

Refs #2064
`temperature_one_support_unchanged` skipped its two routed cases wherever the rejection kernel had no port, so they first ran on ROCm with this change. Their rows are the graph's softmax over the kernel's support (`fused_sample_probs` does not launch the kernel), and gfx1151 lands 2 ulp from the Metal capture at token 26 of the (40, 0.9) row with the support unchanged, in `make verify-rocm` and in three test-fast and one release rerun. The bound goes from 1 to 2 ulp; the support check is unchanged and the stock case keeps its 4.

Refs #2064
From review of #2100: the results page now says what the commits after the measured binary change and quotes `compare_bench_csv.py` (1.024x and 1.130x on the third runs) next to the medians and the per-pair ratios (1.04x, 1.11x, 1.13x); `docs/installation.md` notes the before arm's drift; the upstream README row for LOCAL_FIXES item 30 describes the whole change. `bench_decode.sh` rejects a `--temperature` or `--top-p` value that is not a plain decimal, since both go into the output filename. Two header comments are reflowed.

Refs #2064
Bilingual (en/ko) report on the HIP ports of the Gumbel-max and rejection samplers: the line-for-line ports, the rejection predicate reading the port table, the fast::hip_kernel shape/strides overlay fix (LOCAL_FIXES item 30), the fixed-key, support and chi-square evidence, the gfx1151 measurements, and the 2 ulp snapshot bound.
@inureyes
inureyes force-pushed the update/issue-2064-hip-samplers branch from f931de0 to e2be53a Compare October 4, 2026 18:03
@inureyes
inureyes merged commit 6668031 into main Oct 4, 2026
25 checks passed
@inureyes
inureyes deleted the update/issue-2064-hip-samplers branch October 4, 2026 18:24
inureyes added a commit that referenced this pull request Oct 4, 2026
The ROCm bullet named BitLinear, the SSM update step and the decode-MoE pair; #2100 added the samplers and this branch adds the Mamba1 scan, the fc1 squared-ReLU MoE kernel, xIELU and the add3 LayerNorm.

Refs #2069
inureyes added a commit that referenced this pull request Oct 4, 2026
…2101)

Closes #2069

Four kernels had a Metal port only and sat behind gates that did not read their tables. Each now has a HIP source, a filled `.rocm` entry and a predicate that reads its table, and its gate reads that predicate. CUDA keeps its fallbacks.

## Changes

- **`moe_fc1_relu2`**: HIP port of the fc1 + relu² kernel. `fused_moe_relu2_kernels_available()` gates the `MLXCEL_FUSED_MOE_RELU2` branch of `fused_moe_forward`, whose rows per block default to 2 on ROCm as #2065 measured. The flag stays opt-in.
- **`fused_xielu`**: the early return reads `fused_xielu_kernel_available()` instead of `metal::is_available()`. The HIP kernel follows the ROCm graph's rounding (one rounding to `T` per op, the device `expm1f`, no FP contraction), so it is byte-identical to `apertus_xielu` in f32, f16 and bf16. On ROCm the element count comes from the input shape, so hipRTC compiles one kernel per dtype, not one per length.
- **`mamba1_selective_scan`**: HIP port of the Metal float32-state variant. The CUDA graph-exact variant cannot be reproduced on ROCm (the graph's `state @ C` runs through rocBLAS). `mamba.rs` now reads `mamba1_scan_float_state_kernel_available()` and `mamba1_scan_kernel_accepts`, so Mamba and Falcon-Mamba take the port and a state wider than 32 falls back.
- **`fused_add3_layer_norm`**: `residual_add3_layer_norm` reads `fused_add3_layer_norm_available()` instead of `metal_is_available()`. The HIP kernel reproduces ROCm's own pair (rounded residual adds, the overlay's `layer_norm_kernel<T, 256, 4>`), so the byte-identity test holds unchanged; ROCm builds launch 256 threads per row.
- **CPU device**: these four predicates and #2065's `fused_moe_kernels_available()` / `moe_down_kernel_available()` also require the GPU as the default device, so `MLXCEL_DEVICE=cpu` takes the graph paths instead of throwing "Custom kernels only run on GPU". New test binary `tests/cpu_device_custom_kernel_gates.rs`.
- **Tests**: new `fused_moe_relu2_parity_tests` and `fused_xielu_kernel_matches_graph_every_dtype`; the add3 test module now builds on every backend and adds f32, widths 5 and 1025, and the 6656 limit; the Mamba1 f32 test adds N = 32. No tolerance was loosened.
- **Docs**: technical report (en, ko) under `TECHNICAL_REPORTS/2101-*`; results page `docs/benchmark_results/rocm-metal-only-ports-gfx1151-2026-10-05.md`, raw rows under `benchmarks/`, traces under `benchmarks/logit_traces/rocm_gfx1151_9186b075/`, README, `docs/installation.md` and `docs/environment-variables.md` rows.

`grep -rn '\.rocm = nullptr' --include=*.cpp src/` goes from 16 to 12; the one left in `mlx_cxx_kernels.cpp` is the graph-exact Mamba1 table. A fork bug that faulted the Mamba1 port (by-value `<input>_shape`) was fixed independently by #2100 (LOCAL_FIXES item 30); this branch dropped its own copy of that fix on rebase.

## Measured (gfx1151)

Only Nemotron-H has a checkpoint here; Apertus, Cohere2, Mamba, Falcon-Mamba and Jamba have none and are covered by kernel tests only. Nemotron-3-Nano-30B-A3B-4bit with `MLXCEL_FUSED_MOE_RELU2=1`, `bench_decode.sh` pp512/tg128, before (`main` 6668031, the flag declines to `gather_qmm`) and after alternated, every run through the GPU guard: 75.15 / 75.22 / 75.09 (median 75.15) to 75.24 / 75.38 / 75.13 (median 75.24) tok/s, +0.1%, within noise.

## Deviation from the issue: relu2 greedy text

The issue asked for the relu2 greedy 128-token output to match the default path. It does not: on three prompts the texts agree for about 20 tokens, then part at a near-tie ("user query" against "user request"). The kernel keeps fc1, relu² and fc2 in f32 where `gather_qmm` rounds to bf16 (8.1e-6 nrms from the f32 reference against `gather_qmm`'s 5.1e-3 to 5.4e-3). Teacher-forced `w1ctx512` traces show 0 decided-position mismatches against the ROCm default (1/128 top-1, largest gap 0.5) and against Metal M5 (5/128, largest gap 0.25). The default path is byte-identical to `main`.

## Verification (gfx1151)

- `make verify-rocm` on code head `800e23dc` (rebased on `6668031c`; later commits are docs only): 147 suites, 11867 passed, 0 failed, 378 ignored. `dead_doc_pointers` re-run after the docs commits: passed.
- Parity on gfx1151, each with a negative check that fails when the port is broken (details in the results page): xIELU byte-identical in three dtypes; add3 byte-identical in all ten cases; `mamba1_scan_parity_tests` 5 passed; `fused_moe_relu2_parity_tests` within bounds; `cpu_device_custom_kernel_gates` passes and fails with the device term removed.
- `models::mamba::` and `models::jamba::` 22 passed, 2 ignored; Cohere2 model tests 13 passed, 7 ignored.

Not verified: Metal and CUDA (not available on this host). Metal's kernel sources and table entries are untouched; Metal now runs the widened add3 and Mamba1 tests and the new relu2 test, and its predicates gained the GPU-device term (GPU behavior unchanged; `MLXCEL_DEVICE=cpu` now falls back instead of throwing). CUDA: the fc1_relu2, xIELU, add3 and float-state Mamba1 tables stay null there, and the add3 test compares the unfused pair with itself. Wave64 (CDNA) untested.
inureyes added a commit that referenced this pull request Oct 5, 2026
)

## Summary

On ROCm the three paged-attention port tables had no `.rocm` entry, so the server's batched paged decode, MLA split-KV and the sparse paged decode all took gather-then-SDPA, and 36 tests in the ROCm gate skipped visibly on #1814. This fills the three slots with hipRTC ports of the CUDA bodies and makes the v1 dispatch selector recognise ROCm.

- `paged_attention_hip.h`: the v1 decode, v2 partial and merge CUDA bodies ported to HIP (same names, inputs, outputs, grid and template arguments; `__shfl_xor(v, o, 32)`; `-__builtin_huge_valf()` for `INFINITY`; explicit `(float)` reads for `hip_bfloat16`). Geometry is read from `<input>_shape` as on CUDA, which relies on the overlay fix in #2100. The launches stay in the files that hold the CUDA launches, so the dtype-key pin and its count are unchanged.
- Wave32 guard: v1 and v2 partial fold across 32 lanes and carry the `#error` guard (inert with ROCm 10's clang, as #2067 found). The merge kernel has no lane-level operation (one thread per element, no shuffle, no barrier), so it carries no guard; an `#error` there would only reject a correct kernel on wave64. This deviates from the acceptance checkbox's wording on purpose.
- `PagedV2PartialHolder` and `PagedMergeHolder` take a `GpuKernelBackend` instead of `bool use_cuda`, with `std::call_once` kept.
- The three launchers refuse, on every backend, shapes the bodies cannot index: rank, D of 0, V differing from K in block size, heads or D (axis 0 may differ: MiniMax-M3's sparse launch passes a K allocation with more rows), and a merge head dim outside 1 to 1024.
- `PagedDecodeBackend::Rocm`: returned by `paged_decode_backend()` when the v1 predicate is true and the backend kind is ROCm; follows the CUDA rule in `select_pooled_paged_dispatch`, memo tag 3, and the `gridDim.z <= 65535` guard. Unit tests next to the CUDA ones. The kernel bench example labels ROCm the same way, and the paged autotune ops label ROCm tactics `rocm` instead of `cuda`.
- Stale "no HIP port yet" comments and docs updated (`kernel_ports.rs`, headers, bridge docs, `mla/mod.rs`, Makefile comment, `docs/environment-variables.md`, `docs/installation.md`), plus a results page. The skip macros stay for future ports.

## Measurement

Server paged decode on gfx1151, `scripts/benchmark_paged_decode_production.sh` matrix with Meta-Llama-3.1-8B-Instruct-4bit at `f4ca4b9a` (later commits touch only host checks, comments and the autotune label), fused v2 against `MLXCEL_PAGED_ATTENTION_NATIVE=0` on the same binary, each before/after pair inside one `rocm_gpu_guard.sh` window. Per-request decode tok/s, medians:

| Case | Before | After | |
|---|---|---|---|
| batch 4, ~1K | 5.6 | 6.1 | 1.09x |
| batch 4, ~4K | 10.2 | 10.5 | 1.03x, within spread |
| batch 4, ~16K | 10.9 | 13.9 | 1.28x |
| batch 1, ~16K (1 run) | 17.2 | 24.1 | 1.40x |
| batch 1, ~32K (1 run) | 11.2 | 19.8 | 1.77x |

Batch-4 cases ran three times each (spread at most 0.5 tok/s); the single-sequence cases once, because long guarded windows were rare on the shared host. TTFT is unchanged. The aggregate column of the bench is prefill-dominated at long context and is discussed, not used, in `docs/benchmark_results/rocm-paged-attention-gfx1151-2026-10-05.md`.

## Verification

On gfx1151 (Radeon 8060S, ROCm 10), branch at `f4ca4b9a` (rebased on `c05d5438`):

- `make verify-rocm`: OK. 11869 passed, 0 failed, 378 ignored across 147 test binaries; ROCm smoke OK; no `skipping ... #1814` line (36 on `57d8ed29`).
- After the review fixes (`0d1afcb3`): the paged, MLA, autotune and layers selectors pass (550 tests), `cargo clippy -p mlxcel-core --features rocm --lib --tests -- -D warnings` is clean, `cargo fmt --check` and `dead_doc_pointers` pass. None of the previously skipped tests was edited.
- Mutation check: with the HIP lane fold started at 8 and the HIP merge in base e, 29 of those tests fail (v2 launch, cascade, sparse, split-KV, batched decode, and `test_fused_paged_decode_native_vs_fallback_matrix` for v1). The two v1 tests at head dim 8 do not see the fold change, since 8 dims fit in lanes 0 to 7.
- `make verify-kernel-dtype-keys verify-kernel-port-dispatch`: pass, 9 in scope, `EXPECTED_IN_SCOPE` and its count unchanged.

Not verified: Metal and CUDA (not available on this host). Their kernel bodies, template arguments and `.metal`/`.cuda` entries are unchanged; they are touched by the holder refactor (same `metal_kernel`/`cuda_kernel` calls, one holder per backend) and by the new host shape refusals. No wave64 device was available.

Closes #2068
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:core mlxcel-core: MLX FFI, primitives, KV cache, layers platform:linux Linux (CUDA / packaging) specific priority:low Low priority status:done Completed type:performance Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

perf(rocm): port gumbel_max_sample and rejection_sample to HIP

1 participant