Skip to content

fix(rocm): route large bf16 qmm to hipBLASLt and gate dense prefill - #2085

Merged
inureyes merged 4 commits into
mainfrom
fix/issue-2081-rocm-bf16-dense-gemm
Sep 30, 2026
Merged

inureyes merged 4 commits into
mainfrom
fix/issue-2081-rocm-bf16-dense-gemm

Conversation

@inureyes

@inureyes inureyes commented Sep 30, 2026 •

Copy link
Copy Markdown
Member

Summary

The bf16 case of prefill_dense_gemm_matches_qmm_bytes_where_eligible failed on gfx1151 because quantized_matmul took the fused qmm_wmma_dense_kernel, whose rocWMMA reduction order differs from the hipBLASLt GEMM that dequantize + matmul runs (499 of 1,048,576 outputs, 465 at 1 ULP, 10 at 2, 24 larger, all below 0.01 in magnitude).

Chosen option: 1 (byte identity by routing), with a route check that keeps the guarantee exact on ROCm. Measuring the two kernels showed the WMMA kernel is also the slower one on every bf16 shape from 128 rows up (1.0x to 3.1x at 128 rows, 2.4x to 5.1x at 1024 and 2048), and it wins only on some shapes at 64 rows or fewer. So routing large bf16 GEMMs to dequantize + hipBLASLt makes both sides identical and speeds up prefill. Options 2 and 3 would have kept bf16 prefill on the slow kernel and given up the guarantee.

  • ROCm overlay (LOCAL_FIXES item 29): QuantizedMatmul decides its route in one function, select_qmm_route. On RDNA 3.5 a bf16 GEMM of 128 rows or more goes to dequantize + hipBLASLt. MLX_ROCM_WMMA_QMM_MAX_M overrides the ceiling, and MLX_ROCM_WMMA_QMM=1 removes it (the old dispatch except on low-CU iGPUs). The fp8 route (RDNA 4) and unmeasured tiers keep the kernel.
  • rocm.h exports quantized_matmul_runs_dequant_gemm from that same function. prefill_dense_gemm_eligible asks it through the bridge, so on ROCm the dense path runs only where it returns the qmm bytes: not below the ceiling, not for one row (gemv), not on the fp8 route, not with a batch axis above 1. Metal and CUDA answer true, so their eligibility and the test's assertions there are unchanged.
  • Test: a ROCm-only block checks that any accepted shape matches, and on gfx1151 without MLX_ROCM_* overrides that bf16 at 64 rows (which differs) is refused while dequantize-route shapes are accepted. With the route check bypassed, it fails on bf16 at 64 rows.
  • Production exposure: the dense path only runs on ROCm when MLXCEL_PREFILL_DEQUANT_MIN_M is set. With this change, setting it there changes neither bytes nor speed (Gemma 3 4B 2871 against 2852 tok/s at pp2048).

Prefill on gfx1151, same binary, default against MLX_ROCM_WMMA_QMM=1, three ABBA runs: Gemma 3 4B 4-bit 961 to 2815 tok/s at pp2048 (1146 to 2289 at pp512), Qwen3-0.6B 3059 to 4236 (4581 to 7633), Qwen3-30B-A3B 283 to 297 (311 to 326). Decode is unchanged. Peak memory is within 0.1 GB, except Qwen3-0.6B at pp512 (1.04 to 1.46 GB). Tables: docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md.

Verification

On gfx1151 (Radeon 8060S):

  • cargo test -p mlxcel-core --profile test-fast --features rocm --lib layers::tests::prefill_dense_gemm_matches_qmm_bytes_where_eligible -- --test-threads=1 passes, and also under MLX_ROCM_WMMA_QMM=0, MLX_ROCM_WMMA_QMM=1, MLX_NO_HIPBLASLT=1 and MLX_ROCM_FORCE_LOW_CU=1.
  • cargo clippy -p mlxcel-core --features rocm --lib --tests -- -D warnings passes.
  • make verify-versions verify-kernel-dtype-keys verify-kernel-port-dispatch verify-llama-compat verify-rocm-overlay verify-fmt passes.
  • cargo test --features rocm --test dead_doc_pointers passes.
  • make verify-rocm on 814f9ef: green. 146 test suites, 11,785 passed, 0 failed, 378 ignored; smoke OK (32 tokens). The last commit (c15ff36) only adds a LOCAL_FIXES sentence and two env names to the test's override list; the test, clippy, verify-rocm-overlay and verify-fmt pass on it.

Known limits: the route is judged when the graph is built, and hipBLASLt availability is read again at eval, so a stream capture that starts in between would send both sides to different rocBLAS fallbacks; HIP graphs are off on this backend (use_hip_graphs() returns false), so this needs MLXCEL_PREFILL_DEQUANT_MIN_M set during a decode capture. The dequantize route keeps up to 256 MB of dequantized weights in its LRU after prefill, as it already did for f16.

Not verified: Metal and CUDA are not available on this host. The Rust test and eligibility code changed on those paths, but the new bridge call returns true there and the ROCm test block is gated to ROCm. Other ROCm tiers (RDNA 3, RDNA 4, CDNA) were not measured and keep their previous dispatch.

Closes #2081

On gfx1151 the bf16 case of prefill_dense_gemm_matches_qmm_bytes_where_eligible failed in 499 of 1,048,576 outputs: quantized_matmul took the fused qmm_wmma_dense_kernel, whose rocWMMA reduction order differs from the hipBLASLt GEMM that dequantize + matmul runs. Measuring the two kernels showed the WMMA kernel is also the slower one on every bf16 shape from 128 rows up (1.0x to 5.1x), so the fix makes both sides byte-identical by routing, not by loosening the test.

The ROCm overlay now decides the QuantizedMatmul route in one function, select_qmm_route. On RDNA 3.5 a bf16 GEMM of 128 rows or more goes to dequantize + hipBLASLt (MLX_ROCM_WMMA_QMM_MAX_M overrides, MLX_ROCM_WMMA_QMM=1 restores the old dispatch); the fp8 route and unmeasured tiers keep the kernel (LOCAL_FIXES item 29). rocm.h exports quantized_matmul_runs_dequant_gemm from the same function, and prefill_dense_gemm_eligible asks it through the bridge, so on ROCm the dense path runs only where it returns the qmm bytes. Metal and CUDA answer true and keep the tile rule.

Prefill on gfx1151, default against MLX_ROCM_WMMA_QMM=1: Gemma 3 4B 4-bit 961 to 2815 tok/s at 2048 tokens, Qwen3-0.6B 3059 to 4236, Qwen3-30B-A3B 283 to 297; decode unchanged. The test now also checks, on ROCm, that shapes whose bytes differ are refused; with the route check bypassed it fails on bf16 at 64 rows.

Closes #2081
@inureyes inureyes added status:review Under review type:bug Bug fixes, error corrections, or issue resolutions priority:medium Medium priority area:core mlxcel-core: MLX FFI, primitives, KV cache, layers platform:linux Linux (CUDA / packaging) specific labels Sep 30, 2026
Review follow-ups for #2081. quantized_matmul_runs_dequant_gemm now refuses one-row GEMMs, which matmul sends to gemv rather than hipBLASLt, and takes the device index so the bridge checks the default GPU instead of device 0 (the bridge now compares the device type only). The route-guard test asserts the expected eligibility of each shape and that at least one differing shape is refused on the measured gfx1151 route, runs the route-independent min-rows and narrow-N assertions before any skip, and treats MLX_NO_HIPBLASLT as a route override. Docs say that MLX_ROCM_WMMA_QMM=1 removes the ceiling (the old dispatch except on low-CU iGPUs), when MLX_ROCM_WMMA_QMM_MAX_M applies, and that peak memory for a dense bf16 model above 4B was not measured.

The test passes on gfx1151 by default and under MLX_ROCM_WMMA_QMM=0, MLX_ROCM_WMMA_QMM=1 and MLX_NO_HIPBLASLT=1; clippy on mlxcel-core lib and tests is clean.

Refs #2081
Security review follow-ups for #2081. LOCAL_FIXES item 29 now says that the dequantize route's weight LRU gets no hits across a model's projections in prefill but keeps its last entries (up to 256 MB) alive, as it already did for f16 checkpoints, and that MLX_ROCM_QMM_DEQUANT_CACHE_SIZE=0 turns it off. The ROCm route test also treats MLX_ROCM_FORCE_LOW_CU and MLX_ROCM_FORCE_WARP_SIZE as route overrides, so it skips instead of failing under them.

The test passes on gfx1151 by default and with MLX_ROCM_FORCE_LOW_CU=1; clippy, verify-rocm-overlay and verify-fmt pass.

Refs #2081
@inureyes inureyes added status:done Completed and removed status:review Under review labels Sep 30, 2026
@inureyes
inureyes merged commit f9aefa3 into main Sep 30, 2026
24 checks passed
@inureyes
inureyes deleted the fix/issue-2081-rocm-bf16-dense-gemm branch September 30, 2026 14:29
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:medium Medium priority status:done Completed type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix(rocm): bf16 prefill dense GEMM is not byte-identical to quantized_matmul on gfx1151

1 participant