Repository navigation
update(rocm): route top-k with top-p up to vocab 152064 on ROCm - #2218
Merged
Merged
Conversation
inureyes
force-pushed
the
update/issue-2157-rocm-joint-rejection-cap
branch
from
October 8, 2026 10:40
3f3cb10 to
db16778
Compare
Re-measure the top-k + top-p rejection vocab cap (#2157) on gfx1151 and keep 32768, which is the same crossover M1 Ultra measured. Batch 1 does not clear 1.0 above 32768 (0.98x pipelined at 65536, 0.96x isolated at 128256), so by the issue's rule the ROCm value equals the default and the cap stays one constant with no build-flag branch. - `REJECTION_JOINT_VOCAB_MAX` in `mlx_cxx_bridge.cpp` carries the gfx1151 table in its comment and the rule for moving it per build (`#ifdef MLXCEL_BRIDGE_ROCM_BACKEND`, never a runtime backend check). - New bridge accessor `sampling_rejection_joint_vocab_max()`. `the_joint_vocab_cap_is_the_measured_crossover` pins the value and checks the policy turns exactly there; the routing matrix test adds the 128256 cells. - The not-routed reason string now names the cap and both measurements instead of the M1-only ratio. - `rejection_sampling_microbench` gates on `sampling_rejection_backend_supported()` instead of `custom_kernels_available()`, which is Metal-or-CUDA only and kept the harness from running on ROCm. It adds vocab 128256 and 151936 and prints the build's cap. - `mlxcel-bench-decode --top-k` and `bench_decode.sh --top-k` (filename tag `_k<K>`) so the joint filter can be benchmarked end to end.
…bench - `rejection_sampling_microbench --dtype f32|bf16|f16` casts the synthetic logits, since a bf16 checkpoint hands the sampler bf16 logits and the stock chain's sort cost depends on the dtype; `--config <label>` runs a single filter configuration. - `mlxcel-bench-decode --top-k` rejects negative values; `bench_decode.sh --top-k` rejects leading zeros and values past nine digits so the filename tag matches the value. - Accessor doc comment and the bench script's sampling comment describe the current state.
End-to-end decode on gfx1151 measured the joint top-k + top-p case 1.15x faster on Meta-Llama-3.1-8B-Instruct-4bit (vocab 128256) and 1.12x on Qwen2.5-7B-Instruct-4bit (152064) when routed to the rejection kernel, with every kernel run faster than every chain run, so a ROCm build now caps the joint case at 152064 instead of 32768. The microbenchmark's batch 4 and 8 cells win at every vocabulary; its synthetic batch-1 cell does not clear 1.0 above 32768 in f32 or bf16, and why it does not predict real decode was not isolated. - `REJECTION_JOINT_VOCAB_MAX` in `mlx_cxx_bridge.cpp` is 152064 under `#ifdef MLXCEL_BRIDGE_ROCM_BACKEND` and 32768 otherwise; no runtime backend comparison. Metal and CUDA routing is unchanged. - The not-routed message names the build's cap without quoting a single backend's ratio. - Tests pin the per-build value (`cfg!(feature = "rocm")`), the routing matrix covers 65536 to 152064 per build plus 201088 and 262144 declining everywhere, and the large-vocabulary decline test picks a vocabulary above the build's cap. - Docs: new results page `docs/benchmark_results/rocm-rejection-joint-cap-gfx1151-2026-10-07.md` (f32 and bf16 microbenchmark tables, end-to-end pairs, decision), linked from the samplers page, `installation.md` and `environment-variables.md`. Raw CSVs under `benchmarks/`.
Adds the before/after pairs from this branch's build on current main (Meta-Llama-3.1-8B-Instruct-4bit, `--temperature 0.8 --top-k 40 --top-p 0.95`: 33.15 to 37.73 tok/s, medians of three guarded pairs, both arms' dispatch logged) and quotes those numbers where the docs cite the Llama result.
inureyes
force-pushed
the
update/issue-2157-rocm-joint-rejection-cap
branch
from
October 8, 2026 10:42
db16778 to
7d737d4
Compare
inureyes
added a commit
that referenced
this pull request
Oct 8, 2026
…iling (#2236) ## Summary bf16 affine GEMMs that `select_qmm_route` sends to dequantize + hipBLASLt because they are at or above the WMMA row ceiling (128 rows on RDNA 3.5) no longer go through the dequantized-weight LRU. A forward pass runs more distinct projections than the 8-entry, 256 MB cache holds, so it never hit, and its last entries stayed alive through decode. - New `QmmRoute::DequantGemmAboveWmmaCeiling`, returned only from the ceiling branch. `QuantizedMatmul::eval_gpu` runs the same dequantize and GEMM but into a temporary freed after the GEMM; `quantized_matmul_runs_dequant_gemm` accepts both enumerators. f16 and non-WMMA shapes keep the cache unchanged. - Cache counters (hits, misses, inserts, evictions, bypasses, entries, bytes) as relaxed atomics: `rocm::dequant_cache_stats()` in the overlay (stubbed in `no_rocm.cpp`), `mlxcel_core::rocm_qmm_cache::dequant_cache_stats()` through the bridge, and `MLX_ROCM_QMM_DEQUANT_CACHE_STATS=1` to print them at exit. - `tests/rocm_qmm_dequant_cache.rs`, LOCAL_FIXES item 29, `docs/environment-variables.md`, and a dated section in `docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md`. ## Measurements (gfx1151, pp2048/tg128, medians of 3 on 8753883, each run under `rocm_gpu_guard.sh --idle-secs 60`) | Model | Arm | Prefill tok/s | Decode tok/s | MLX peak | Active after prefill | Cache hits / misses / bypassed, entries at exit | |---|---|---:|---:|---:|---:|---| | gemma-3-4b-it-4bit (bf16) | before | 2872 (2819-2910) | 62.75 (62.68-62.76) | 4.52 GB | 2.80 GB | 0 / 340 / 0, 6 (230 MiB) | | gemma-3-4b-it-4bit (bf16) | after | 2868 (2839-2945) | 62.71 (62.55-62.73) | 4.52 GB | 2.56 GB | 0 / 0 / 340, 0 | | Llama-3.1-8B-Instruct-4bit (f16) | before | 1136 (1106-1136) | 36.63 (32.36-36.84) | 6.92 GB | 4.75 GB | 0 / 320 / 0, 2 (224 MiB) | | Llama-3.1-8B-Instruct-4bit (f16) | after | 1132 (1048-1143) | 36.67 (36.58-36.92) | 6.92 GB | 4.75 GB | 0 / 320 / 0, 2 (224 MiB) | Prefill and decode differences are inside the run-to-run range; no speedup is claimed. Before counts come from a counters-only build that routes like 8753883. One confirmation run per model on the rebased head (main ad84435) matched the after rows. The f16 cache also never hits, so #2232 tracks turning its default off. ## Verification All on gfx1151 (Radeon 8060S), ROCm 10.0.0 / HIP 7.15. The two test binaries and the overlay, dtype-key, port-dispatch and fmt gates were rerun on the final head (rebased onto main ae343d8): - `cargo test --release --features rocm --test rocm_qmm_dequant_cache -- --test-threads=1`: pass. With the ceiling branch returning `DequantGemm` again, the two bf16 ceiling cases fail (`misses=1 inserts=1 bypasses=0 entries=1` on the first pass); the forced-WMMA and f16 cases still pass. - `cargo test --release --features rocm --test rocm_qmm_env -- --test-threads=1`: pass. - Logit traces (`examples/logit_trace`, Gemma 3 4B, 4 x 256-token chunks) before and after: byte-identical; `scripts/compare_logit_traces.py --decided 2.0` reports 0 / 1024 disagreements. - `make verify-rocm-overlay verify-versions verify-kernel-dtype-keys verify-kernel-port-dispatch verify-llama-compat verify-fmt`, `cargo test --test dead_doc_pointers`, `cargo clippy -p mlxcel-core --features rocm --lib --tests -- -D warnings`, `cargo clippy -p mlxcel --features rocm --test rocm_qmm_dequant_cache -- -D warnings` (before the rebase): pass. - Full `make verify-rocm` with `MLXCEL_ROCM_SMOKE_MODEL=models/mlx/Qwen3-0.6B-4bit` at `79a8c9bb` on `6dbfe7d7` (the later commit only adds the technical report): `[verify-rocm] OK`, 162 cargo test suites, 12191 passed, 0 failed, 399 ignored; the ROCm smoke generated 32 tokens on the GPU. An earlier run on the `ae343d84` base was also green (12190 passed, 0 failed) before main gained #2218, which touches the bridge files. - Review: the orchestrator read the route, cache and bridge diff (relaxed atomic counters bumped once per GEMM on the dequant routes only, the exit report registered after the cache is constructed, the bridge returns zeros without `MLXCEL_BRIDGE_ROCM_BACKEND`); no CRITICAL or HIGH findings. Not verified: Metal and CUDA (not available on this host). The change is in the ROCm overlay and the bridge function is stubbed to zeros off ROCm (`#ifdef MLXCEL_BRIDGE_ROCM_BACKEND`); the unit test `stats_are_zero_off_rocm` covers that stub but has not been built or run on those backends. Closes #2151
3 tasks done
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Re-measures the joint top-k + top-p rejection vocab cap (
REJECTION_JOINT_VOCAB_MAX) on gfx1151. A ROCm build now routes top-k with top-p to the rejection kernel up to vocab 152064; Metal and CUDA keep 32768. With--temperature 0.8 --top-k 40 --top-p 0.95, Llama 3.1 8B decode on this branch goes from 33.15 to 37.73 tok/s (1.14x).Where the cap lives (for epic #2166 Phase 2):
REJECTION_JOINT_VOCAB_MAXinsrc/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp, read only bysampling_rejection_routesand reported by the new bridge accessorsampling_rejection_joint_vocab_max(). A unified per-row sampling step should callsampling_rejection_routes(or the accessor) rather than copy the number. The per-build value is pinned bythe_joint_vocab_cap_is_the_measured_crossoverinsrc/lib/mlxcel-core/src/sampling_rejection_tests.rs.Decision and evidence
Full tables:
docs/benchmark_results/rocm-rejection-joint-cap-gfx1151-2026-10-07.md.rejection_sampling_microbench, three guarded runs, a measurement build with the cap lifted so every joint cell routes): batches 4 and 8 win at every vocabulary (pipelined 1.17x to 2.02x), but the synthetic batch-1 cell does not clear 1.0 above 32768 (0.90x to 0.98x pipelined, 0.92x to 0.96x isolated at 128256 and up). A bf16 rerun reads the same, so the logits dtype is not the cause.Changes
mlx_cxx_bridge.cpp:REJECTION_JOINT_VOCAB_MAXis 152064 under#ifdef MLXCEL_BRIDGE_ROCM_BACKENDand 32768 otherwise. This uses the build-flag pattern of the fused-MoE SGY default (update(rocm): port the fused MoE decode kernels to HIP #2098) and thecfg!(feature = "rocm")defaults of update(rocm): enable fused add-RMSNorm and RoPE-append by default on ROCm #2189, with no runtime backend comparison, soverify-kernel-port-dispatchpasses. The gfx1151 tables are in the comment above the constant. The not-routed message names the build's cap.sampling_rejection_joint_vocab_max().examples/rejection_sampling_microbench.rs: adds vocab 128256 and 151936 and the--dtype f32|bf16|f16and--config <label>options, and prints the build's cap. It also fixes the port gate: it checkedcustom_kernels_available(), which is Metal-or-CUDA only, so it refused to run on ROCm. It now checkssampling_rejection_backend_supported().mlxcel-bench-decode --top-kandscripts/bench_decode.sh --top-k(filename tag_k<K>, validated), so the joint filter can be benchmarked end to end.rocm-samplers-gfx1151-2026-10-05.md. The ROCm sampled-decode row ininstallation.mdnow lists which filter combinations route at the Llama 3 and Qwen vocabularies, and theMLXCEL_SAMPLING_REJECTIONrow inenvironment-variables.mdis updated. Raw CSVs are underbenchmarks/.Verification
All on gfx1151, branch rebased on
mainatad844354:make -k verify-rocm: OK. 12187 passed, 0 failed, 398 ignored across 161 test binaries; fmt, clippy (-D warnings), the overlay and kernel-dispatch checks, and the smoke run pass.cargo test --release --features rocm -p mlxcel-core --lib -- --test-threads=1 sampling_rejection_tests sampling_fixed_key_tests: 34 passed, covering the support, distribution and fixed-key tests.cargo test --release --features rocm --test sampling_gumbel_kill_switch --test sampling_rejection_kill_switchpass.cargo test --test dead_doc_pointerspasses.mlxcel generatelogsrejection kernel (batch 1, vocab 128256, ... top_k 40, top_p 0.95 ...)by default andargpartition chain: pinned by MLXCEL_SAMPLING_REJECTIONwith the switch set.scripts/rocm_gpu_guard.sh. Every accepted attempt was clean, and one microbenchmark run was retried after a rejection.Not verified: Metal and CUDA are not available on this host. Their cap value (32768) is unchanged and the routing code is shared, but the non-ROCm arm of the per-build tests was not run here. Qwen2.5-7B was measured on the measurement build only, not re-run on the rebased build.
Closes #2157