diff --git a/TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.en.md b/TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.en.md new file mode 100644 index 000000000..b0a1e6ab8 --- /dev/null +++ b/TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.en.md @@ -0,0 +1,203 @@ +# Technical Report: PR #2085 - Route large bf16 qmm to hipBLASLt and gate dense prefill on ROCm + +**Date**: 2026-09-30 + +**Status**: Implemented and validated on the gfx1151 host; head `c15ff36d` on origin/main `7278397a`, pending merge. + +**Languages**: C++/HIP (ROCm overlay: `QuantizedMatmul` dispatch, exported route predicate), C++ (cxx bridge), Rust (dense-prefill eligibility, test), Markdown + +**Risk Level**: Medium (the default kernel for every bf16 affine GEMM of 128 rows or more changes on the RDNA 3.5 tier, which covers prefill of every bf16-scale 4-bit checkpoint there; other ROCm tiers keep their dispatch unless an environment variable sets a ceiling, and on Metal and CUDA the new bridge call answers true, so their eligibility is unchanged) + +## Executive Summary + +Issue #2081 (part of epic #1801) was the last failure in `make verify-rocm` on gfx1151: the bf16 case of `layers::tests::prefill_dense_gemm_matches_qmm_bytes_where_eligible`. The test asserts that mlxcel's dense prefill path (`dequantize` then `matmul`) returns the same bytes as `quantized_matmul` wherever `prefill_dense_gemm_eligible` accepts a projection. On ROCm, bf16 `quantized_matmul` ran the fork's fused `qmm_wmma_dense_kernel`, which accumulates through rocWMMA tiles in its own K order, while the dense side ran hipBLASLt. 499 of 1,048,576 outputs differed. f16 passed because its `quantized_matmul` also dequantizes and calls hipBLASLt. + +The issue offered three options: make the two sides byte-identical, exclude the ROCm bf16 case from eligibility, or replace byte identity with a ULP bound. The PR takes option 1, and does it by routing rather than by changing either kernel. An op-level measurement showed that the WMMA kernel is also the slower of the two on every bf16 shape measured from 128 rows up, so sending those GEMMs to dequantize + hipBLASLt makes both sides identical and speeds up production prefill. With the default route against the old one (`MLX_ROCM_WMMA_QMM=1`), Gemma 3 4B 4-bit prefill goes from 961 to 2815 tok/s at 2048 tokens. + +The dispatch decision lives in one overlay function, `select_qmm_route`, which feeds both `QuantizedMatmul::eval_gpu` and an exported predicate, `rocm::quantized_matmul_runs_dequant_gemm`. The Rust eligibility asks that predicate through the bridge, so on ROCm the dense path runs only where `quantized_matmul` already runs the same GEMM. The unit's full `make verify-rocm` on `814f9efb` was the first fully green ROCm gate of the epic #1801 run. + +## 1. Problem Statement + +### 1.1 The failing assertion + +``` +panicked at src/lib/mlxcel-core/src/layers.rs:6369:17: +assertion `left == right` failed: dtype 12 bias false: dense GEMM must match qmm bytes +``` + +The issue measured the first failing case (bf16, no bias, x `[1, 1024, 2048]`, 4-bit affine weight `[1024, 2048]`, group 64): 499 of 1,048,576 outputs differ (0.048%). 465 are 1 ULP, 10 are 2 ULP, and 24 are 3 to 35 ULP, all on outputs below 0.01 in magnitude, where cancellation makes bf16 ULP distance large. The largest absolute error is 0.25 (41.5 against 41.75, 1 ULP) against a maximum |out| of 76.5, and one near-zero element crosses sign (-8.5e-6 against 1.6e-5, 28308 ULP by bit distance). The bias case was never reached. The failure had been hidden behind the NVFP4 abort that #1806 (PR #2030) removed, and had failed in every full ROCm gate since. + +### 1.2 Why the premise held on Metal and not on ROCm + +The eligibility rule came from #1994/#2001 (PR #2002): affine mode, f16 or bf16 input with scales in the same dtype, a 2-D weight, at least `min_rows` rows, and more than 512 output tiles of 32 x 32. Its byte-identity claim rests on an in-tree sweep that ran on Metal, where both sides dequantize with the same rounding and differ only in tiling. On ROCm the two sides reach different kernels: + +- Dense side: `affine_dequantize`, then `matmul`, which runs hipBLASLt. +- qmm side, bf16: `qmm_wmma_dense_kernel` whenever the device has native WMMA and is not a low-CU iGPU, for bf16 x/scales/biases, group 64, 4/6/8 bits, `N % 16 == 0`, `K % 64 == 0`. Its dequantization matches `affine_dequantize`, but it accumulates in f32 through rocWMMA 16 x 16 x 16 tiles in its own K order. +- qmm side, f16: no WMMA path, so `affine_dequantize` plus `dequant_rocblas_gemm`, which is hipBLASLt too. + +`MLX_ROCM_WMMA_QMM=0` made the test pass, which confirmed a reduction-order difference and not a dequantization one. The issue also ruled out the CPU-stream OpenBLAS miswrite fixed in #2079: every op in the test runs on the default GPU stream. + +### 1.3 Production exposure before the fix + +`prefill_dense_gemm_min_rows_default` in `hardware.rs` returns a threshold only on Apple M1. On ROCm the dense path runs only when an operator sets `MLXCEL_PREFILL_DEQUANT_MIN_M`, so the mismatch was a latent correctness hole behind an opt-in, not a default-path bug. + +## 2. Why Option 1, by Routing + +The issue required a decision among three options and forbade simply `cfg`-gating the assertion off. It also attached a condition to option 1: a fix is acceptable only if it does not regress bf16 prefill throughput on gfx1151, compared at the op level and on a real model. + +The unit measured the two kernels before choosing. bf16, 4-bit g64, one `[1, M, K]` input against an `[N, K]` weight, mean of 40 calls per arm in two alternated rounds (`dense` re-dequantizes on every call): + +| M | K x N | qmm ms | dense ms | dense / qmm | differ | +|---:|---|---:|---:|---:|---:| +| 16 | 4096 x 4096 | 0.388 | 0.337 | 0.87 | 79 / 65,536 | +| 16 | 4096 x 14336 | 1.329 | 1.649 | 1.24 | 261 / 229,376 | +| 64 | 4096 x 4096 | 0.413 | 0.488 | 1.18 | 305 / 262,144 | +| 64 | 14336 x 4096 | 1.472 | 1.995 | 1.36 | 1,217 / 262,144 | +| 128 | 4096 x 4096 | 0.805 | 0.570 | 0.71 | 572 / 524,288 | +| 128 | 14336 x 4096 | 2.519 | 2.479 | 0.98 | 2,309 / 524,288 | +| 256 | 4096 x 14336 | 4.022 | 2.274 | 0.57 | 4,153 / 3,670,016 | +| 512 | 4096 x 14336 | 9.187 | 3.769 | 0.41 | 4,242 / 7,340,032 | +| 1024 | 4096 x 1024 | 1.403 | 0.320 | 0.23 | 1,164 / 1,048,576 | +| 1024 | 4096 x 14336 | 15.775 | 4.644 | 0.29 | 12,251 / 14,680,064 | +| 2048 | 14336 x 4096 | 51.491 | 10.046 | 0.20 | 29,524 / 8,388,608 | + +Across the full sweep of 50 cells (M 2 to 2048; K x N of 4096 x 4096, 4096 x 1024, 4096 x 14336, 14336 x 4096, 1024 x 3072), the WMMA kernel won only on some shapes at 64 rows or fewer, and dense was faster on every shape from 128 rows up (1.0x to 3.1x at 128 rows, 2.4x to 5.1x at 1024 and 2048, per the PR). Every cell also differed in bytes, so the mismatch is not specific to the test shape. + +That measurement settled the choice, because correctness and speed pointed the same way: + +- **Option 1 by routing** sends large bf16 GEMMs to the route that is both faster and byte-identical to the dense path. With the change, `quantized_matmul` differs from dense in 0 outputs of every cell from 128 rows up, and below 128 rows it is unchanged. +- **Option 1 by changing the WMMA kernel** (matching hipBLASLt's reduction order) would have kept the slower kernel on exactly the shapes where it loses, and it would have tied the kernel to hipBLASLt's internal reduction order, which is not a stable target across library versions. +- **Option 2** (exclude ROCm bf16 on WMMA devices from eligibility) would have made the test honest but left bf16 prefill on the slow kernel, and the opt-in dense path would never help a bf16 model on ROCm. +- **Option 3** (a ULP bound) would have given up the guarantee the dense path is built on, and the measured distribution makes a clean bound hard to state: 24 outputs beyond 2 ULP and one sign crossing at 28308 ULP bit distance mean the bound would need an absolute term fitted to near-zero outputs. The issue specifically warned against a bound chosen to make the test pass. It would also have kept bf16 prefill on the slow kernel. + +Below 128 rows the kernel still wins on some shapes, so it stays there, and the dense path is refused there instead (section 3.3). That part is option 2 applied only where routing does not reach. + +## 3. Change Summary + +Three commits on `fix/issue-2081-rocm-bf16-dense-gemm`: + +- **`a73370f2`** `fix(rocm): route large bf16 qmm to hipBLASLt and gate dense prefill`: the route function, the ceiling, the exported predicate, the bridge call, the eligibility change, the ROCm test block, the results page and the docs. +- **`814f9efb`** `fix(rocm): tighten the dense-prefill route check and its test`: review follow-ups. The predicate refuses one-row GEMMs (which `matmul` sends to gemv, not hipBLASLt) and takes a device index so the bridge checks the default GPU instead of device 0. The route-guard test asserts the expected eligibility of each shape and that at least one differing shape is refused, runs the route-independent assertions before any skip, and treats `MLX_NO_HIPBLASLT` as a route override. +- **`c15ff36d`** `docs(rocm): note the dequant cache footprint and more route overrides`: security-review follow-ups. `LOCAL_FIXES.md` item 29 documents the dequantized-weight LRU footprint, and the test's override list gains `MLX_ROCM_FORCE_LOW_CU` and `MLX_ROCM_FORCE_WARP_SIZE`. + +Files by area: + +- Overlay: `patches-rocm/mlx/backend/rocm/quantized/qmm.hip` (`QmmRoute`, `QmmRouteInputs`, `wmma_qmm_env`, `wmma_qmm_max_m`, `select_qmm_route`, `quantized_matmul_runs_dequant_gemm`, dispatch rewired), `rocm.h` (declaration), `no_rocm.cpp` (stub returning false), `LOCAL_FIXES.md` item 29. +- Bridge: `src/lib/mlxcel-core/cpp/mlx_cxx_bridge.{h,cpp}` (`quantized_matmul_matches_dense_gemm`), `src/lib/mlxcel-core/src/lib.rs` (ffi declaration). +- Rust: `src/lib/mlxcel-core/src/layers.rs` (eligibility, doc comment, test), `src/lib/mlxcel-core/src/hardware.rs` (doc comment only). +- Docs: `docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md`, `docs/environment-variables.md`, `docs/mlxcelverse/upstream/README.md` (item 29 not packaged yet). + +### 3.1 One route decision + +Before the PR, `QuantizedMatmul::eval_gpu` decided its route inline across three separate conditions: a WMMA block, a dequant-GEMM block, and inside it an fp8 block. The PR moves that into `select_qmm_route(const QmmRouteInputs&, rocm::Device&)`, which returns one of `WmmaDense`, `DequantFp8Gemm`, `DequantGemm` or `Other`, and `eval_gpu` now branches on the returned value. The WMMA shape test is the same conjunction as before. The change in behavior is one condition inside the WMMA branch: + +```cpp +const bool hipblaslt_is_faster = + env != 1 && dequant && in.M >= wmma_qmm_max_m(d) && !fp8(); +``` + +When it holds, a GEMM that fits the WMMA kernel goes to `DequantGemm` instead. `wmma_qmm_max_m` returns `MLX_ROCM_WMMA_QMM_MAX_M` when that is a positive integer, else 128 on the `Rdna35` tier and `INT_MAX` everywhere else. Three guards keep the ceiling from reaching a route it was not measured against: + +- `dequant` must hold: the dequant route must be available and enabled (`MLX_ROCM_QMM_DEQUANT_GEMM` not `0`) and preferred or forced for the shape. The ceiling never sends a GEMM to `Other`. +- `!fp8()`: where the fallback would be the e4m3 path (RDNA 4 with an fp8-capable hipBLASLt), the fused kernel stays, as before. The fp8 lambda is evaluated only where it decides the route, so a device that never reaches it is not probed. +- `env != 1`: `MLX_ROCM_WMMA_QMM=1` removes the ceiling, which is the old dispatch on any device that is not a low-CU iGPU (there it forces the kernel on, as it always did). + +### 3.2 The exported predicate + +`rocm::quantized_matmul_runs_dequant_gemm(device_index, M, N, K, x_dtype, scales_dtype, biases_dtype, group_size, bits)` rebuilds the `QmmRouteInputs` that `eval_gpu` derives for one transposed affine GEMM with no batch dimensions, including `should_use_dequant_gemm_path` and the `force_dequant_gemm` term for bit widths qmv does not support, and calls the same `select_qmm_route`. It returns true only for `QmmRoute::DequantGemm` with `is_hipblaslt_available()`: + +- `M < 2` is refused because `matmul` sends a single row with a transposed weight to gemv, not hipBLASLt. +- `DequantFp8Gemm` is refused because the fp8 GEMM rounds the activation and weight to e4m3 and cannot match a bf16 matmul. +- Without hipBLASLt, `dequant_rocblas_gemm` and `matmul`'s rocBLAS path fall back to different rocBLAS calls, so only the hipBLASLt case is claimed. + +Because both callers go through the same function, the predicate cannot drift from the dispatch as long as the inputs are rebuilt the same way. The environment overrides (`MLX_ROCM_WMMA_QMM`, `MLX_ROCM_WMMA_QMM_MAX_M`, `MLX_ROCM_QMM_DEQUANT_GEMM` and the rest) move both answers together. + +### 3.3 The bridge and the Rust eligibility + +`quantized_matmul_matches_dense_gemm` in the bridge answers true unless the build is a ROCm bridge (`MLXCEL_BRIDGE_ROCM_BACKEND`), the runtime backend is ROCm and the default device is a GPU. On ROCm it refuses any x with an axis above 1 before the last two, because `QuantizedMatmul` batches over those axes and a batched GEMM is not the single GEMM `matmul` runs on the same rows, then asks the overlay predicate with the default GPU's index. `no_rocm.cpp` stubs the predicate as false for fork builds without ROCm, and those builds never reach it because of the backend check. + +`prefill_dense_gemm_eligible` keeps its tile rule and adds the bridge call as the last condition, after the cheaper row and tile checks return early. On Metal and CUDA the call returns true, so eligibility there is exactly the old rule. The doc comment now carries the #2081 evidence (mismatch count, ULP distribution, kernel paths) that the acceptance criteria required, and states the invariant the change enforces: the dense path is an optimization, so wherever it runs it must return the bytes `quantized_matmul` would have. + +### 3.4 The test + +The Metal-era test body is unchanged in intent. The route-independent assertions (the narrow-N refusal at exactly 512 tiles, and the min-rows refusal) now run first. The byte check at `[1, 1024, 2048]` still asserts eligibility for both dtypes, except that on a ROCm device other than the measured gfx1151 route, a refused shape logs and skips instead of failing. + +A ROCm-only helper, `prefill_dense_gemm_rocm_route_guard`, checks the guarantee itself on shapes that pass the tile rule but reach different routes: + +| dtype | rows | N | route on gfx1151 | expected eligibility | +|---|---:|---:|---|---| +| bf16 | 64 | 8448 | WMMA kernel (below the ceiling) | refused | +| f16 | 64 | 8448 | dequantize + hipBLASLt | accepted | +| bf16 | 256 | 4096 | dequantize + hipBLASLt (above the ceiling) | accepted | + +N 8448 at 64 rows gives 2 x 264 = 528 tiles, just above the 512-tile floor, so the tile rule does not hide the route check. On every ROCm device the helper asserts that an accepted shape matches in bytes. On the measured route (`rocm_measured_route()`: gfx1151 and none of seven route-moving variables set) it also asserts each expected eligibility and that at least one shape really differed, so the guard cannot pass with the route check removed. With the route check bypassed, it fails on bf16 at 64 rows. Finally, x `[2, 512, 2048]` must be refused on ROCm because of its batch axis. + +## 4. Results + +### 4.1 Model prefill + +`mlxcel-bench-decode --prompt-tokens {512, 2048} -n 8 --warmup-tokens 4 --ignore-eos` on gfx1151, one binary, `MLX_ROCM_WMMA_QMM=1` ("before", the old dispatch on this device) against the default ("after"), three ABBA runs per cell, with a sampler confirming no other GPU process ran. Prefill tok/s, mean (range): + +| Model | Scales | pp512 before | pp512 after | pp2048 before | pp2048 after | +|---|---|---:|---:|---:|---:| +| gemma-3-4b-it-4bit | bf16 | 1146 (1046-1199) | 2289 (2197-2359) | 961 (948-972) | 2815 (2805-2824) | +| Qwen3-0.6B-4bit | bf16 | 4581 (4230-5024) | 7633 (7147-8012) | 3059 (2894-3219) | 4236 (4027-4566) | +| Qwen3-30B-A3B-4bit | bf16 | 311 (305-315) | 326 (321-334) | 283 (282-284) | 297 (296-299) | +| Meta-Llama-3.1-8B-Instruct-4bit | f16 | 969 (879-1032) | 1008 (994-1026) | 1149 (1143-1153) | 1132 (1124-1137) | + +Computed from the means: Gemma 3 4B gains 2.0x at 512 tokens and 2.9x at 2048; Qwen3-0.6B gains 67% and 38%; Qwen3-30B-A3B gains about 5% at both lengths, because only its attention projections take this path and its experts go through `gather_qmm`. + +Llama 3.1 8B is the control. Its scales are f16, so none of its GEMMs reaches the changed branch (f16 never takes the WMMA kernel), and its movement (+4.0% at 512 with overlapping ranges, -1.5% at 2048) has no code path to explain it. It bounds the run-to-run noise the bf16 gains should be read against. + +Decode is unchanged in every cell: M is 1, which never took the WMMA kernel. MLX peak memory is within 0.1 GB of before, except Qwen3-0.6B at 512 tokens (1.04 to 1.46 GB), because the dequantize route allocates a bf16 copy of each weight matrix, which `LOCAL_FIXES.md` item 28 (PR #2084) bounds. + +### 4.2 Production exposure: `MLXCEL_PREFILL_DEQUANT_MIN_M` + +On ROCm the dense path runs only when `MLXCEL_PREFILL_DEQUANT_MIN_M` is set. After this change, the only projections it accepts there are ones where `quantized_matmul` already runs dequantize + hipBLASLt, so turning it on cannot change bytes. It also cannot change speed in any meaningful way: both paths run the same GEMM, and `quantized_matmul` does it with its dequantized-weight cache. Measured with `MLXCEL_PREFILL_DEQUANT_MIN_M=1024` against unset, same binary, pp2048, three ABBA runs: gemma-3-4b-it-4bit 2871 against 2852 tok/s, Qwen3-0.6B-4bit 4174 against 4185. + +So the production picture is: + +- **Default configuration (variable unset)**: the dense path never runs on ROCm, before or after. The user-visible effect of the PR is the faster `quantized_matmul` route for bf16 prefill. +- **Variable set**: before the PR, an operator could get bf16 outputs different from `quantized_matmul` (the #2081 mismatch) on any accepted projection. After it, the variable is a no-op for both bytes and speed. `docs/environment-variables.md` and the `hardware.rs` doc comment now say that setting it on ROCm gains nothing, and the ROCm default stays off. + +## 5. Technical Decisions + +- **Measure the kernels before choosing between the three options.** The issue's throughput condition on option 1 turned into the deciding evidence: the kernel that broke byte identity was also the slower one from 128 rows up. +- **Route rather than rewrite a kernel.** Moving a shape between two existing, already-shipped routes needs no new numerical code and makes identity hold by construction (the same GEMM on both sides), instead of by matching an opaque library's reduction order. +- **One function for dispatch and predicate.** A separate predicate that restated the dispatch conditions would drift the first time someone edited one side. Sharing `select_qmm_route` makes the Rust eligibility a question to the dispatch itself. +- **Claim only the hipBLASLt case.** The predicate refuses one-row GEMMs, the fp8 route, the no-hipBLASLt fallback, and batched inputs, because in each of those the two sides reach different code even when the high-level route name matches. +- **Scope the ceiling to the measured tier.** The 128-row default applies only to `Rdna35`. RDNA 3, RDNA 4 and CDNA keep the kernel at every row count unless `MLX_ROCM_WMMA_QMM_MAX_M` is set, and the fp8 fallback is excluded, so an unmeasured device cannot lose throughput by default. +- **Keep an exact switch for the old dispatch.** `MLX_ROCM_WMMA_QMM=1` restores the pre-PR route on this device, which is how one binary produced both columns of the prefill table. +- **Leave the ROCm default of `MLXCEL_PREFILL_DEQUANT_MIN_M` off.** With identical routes the dense path has nothing to add on ROCm, as 4.2 measured. + +## 6. Validation + +From the PR body, on gfx1151 (Radeon 8060S): + +- The target test passes by default and 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, and `cargo test --features rocm --test dead_doc_pointers` passes. + +Orchestrator verification: + +- The unit ran the full `make verify-rocm` on `814f9efb` (gfx1151): 146 test suites, 11,785 passed, 0 failed, 378 ignored, smoke OK. That is the first fully green ROCm gate of the epic #1801 run; the previous baseline failure was this test (PR #2084's gate on `3c9edea0` failed in exactly this target and nowhere else). +- The head commit `c15ff36d` only adds a `LOCAL_FIXES.md` sentence and two environment names to the test's override list. The target test, clippy, `verify-rocm-overlay` and `verify-fmt` passed on it. +- The orchestrator runs a final gate on a fresh build after merge. + +## 7. Residual Risks and What Was Not Verified + +- **The route is judged at graph build.** `prefill_dense_gemm_eligible` asks the predicate when the graph is built, and `dequant_rocblas_gemm` reads hipBLASLt availability again at eval. A stream capture starting 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 predicate also assumes the hipBLASLt launch itself does not throw; if it did, each side would take its own rocBLAS fallback. +- **Dequantized-weight cache footprint.** The route's LRU (8 matrices or 256 MB, `MLX_ROCM_QMM_DEQUANT_CACHE_SIZE`, `MLX_ROCM_QMM_DEQUANT_CACHE_MAX_BYTES`) gets no hits across a model's projections in prefill, so each projection is dequantized again per prefill chunk, and the last entries (up to 256 MB) stay alive after prefill. f16 checkpoints already behaved this way; bf16 checkpoints on RDNA 3.5 now do too. `MLX_ROCM_QMM_DEQUANT_CACHE_SIZE=0` turns it off. The transient copies grow with projection size; no dense bf16-scale checkpoint above 4B was available on the host, so peak memory for one was not measured. +- **Other RDNA tiers were not measured.** The ceiling applies to the whole `Rdna35` tier (`gfx1150` to `gfx1152`) but was measured on gfx1151 only; low-CU gfx1152 parts skip the WMMA kernel anyway unless forced. RDNA 3, RDNA 4 and CDNA keep their previous dispatch. On a ROCm device other than gfx1151, the test's byte check skips a refused shape instead of failing, so coverage there is weaker. +- **Rows between 64 and 128.** The published table has no cell between those row counts; 128 is the first measured count at which dense won on every shape. +- **Metal and CUDA were not run on this host.** The Rust eligibility and test changed on those paths, but the bridge returns true there and the ROCm block is gated to ROCm. +- **Upstreaming.** Item 29 is not packaged for the fork: the ceiling was measured on one device, and the exported predicate exists for mlxcel's dense-prefill check. A fork PR would carry the ceiling alone, with measurements from at least one more RDNA 3 or RDNA 3.5 device. + +## 8. Learning Points + +- **A byte-identity claim is a claim about kernels, not about math.** The #2001 sweep was valid on Metal. On another backend the same two ops reached different kernels, and the premise needed a backend-level check rather than a tile rule. +- **When a correctness fix and a performance question share a knob, measure both first.** Here the evidence that fixed the test also closed a 2.9x prefill gap the issue had not asked about. +- **Make a predicate about dispatch ask the dispatch.** Exporting a function that runs the same route selection, instead of mirroring conditions in Rust, is what keeps the eligibility honest when environment overrides or future tiers move the route. +- **Give a guard test a mutation it must catch.** The route-guard asserts that at least one shape really differs on the measured route, so deleting the route check fails the test instead of silently passing it. + +Refs: #2081, #1801, #1994, #2001, #2002, #1806, #2030, #2079, #2084. diff --git a/TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.ko.md b/TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.ko.md new file mode 100644 index 000000000..2b944f5a2 --- /dev/null +++ b/TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.ko.md @@ -0,0 +1,203 @@ +# 기술 보고서: PR #2085 - ROCm에서 큰 bf16 qmm을 hipBLASLt로 보내고 dense prefill을 경로로 제한 + +**날짜**: 2026-09-30 + +**상태**: gfx1151 호스트에서 구현 및 검증 완료. head `c15ff36d`, origin/main `7278397a` 기준, 머지 대기 중. + +**언어**: C++/HIP (ROCm overlay: `QuantizedMatmul` dispatch, export된 경로 predicate), C++ (cxx bridge), Rust (dense prefill 적격성, 테스트), Markdown + +**위험도**: 중간 (RDNA 3.5 tier에서 128 row 이상인 모든 bf16 affine GEMM의 기본 kernel이 바뀌며, 이는 그 tier에서 bf16 scale 4-bit checkpoint의 prefill 전체에 해당합니다. 다른 ROCm tier는 환경 변수로 ceiling을 지정하지 않는 한 기존 dispatch를 유지하고, Metal과 CUDA에서는 새 bridge 호출이 true를 반환하므로 적격성이 바뀌지 않습니다) + +## 요약 + +이슈 #2081(epic #1801의 일부)은 gfx1151의 `make verify-rocm`에 남은 마지막 실패였습니다. `layers::tests::prefill_dense_gemm_matches_qmm_bytes_where_eligible`의 bf16 케이스입니다. 이 테스트는 `prefill_dense_gemm_eligible`이 projection을 받아들이는 곳마다 mlxcel의 dense prefill 경로(`dequantize` 후 `matmul`)가 `quantized_matmul`과 같은 바이트를 반환하는지 확인합니다. ROCm에서 bf16 `quantized_matmul`은 fork의 fused `qmm_wmma_dense_kernel`을 실행했고, 이 kernel은 rocWMMA tile을 통해 자체 K 순서로 누적합니다. 반면 dense 쪽은 hipBLASLt를 실행했습니다. 1,048,576개 출력 중 499개가 달랐습니다. f16은 `quantized_matmul`도 dequantize 후 hipBLASLt를 호출하기 때문에 통과했습니다. + +이슈는 세 가지 선택지를 제시했습니다. 두 경로를 바이트 단위로 같게 만들기, ROCm bf16 케이스를 적격성에서 제외하기, 바이트 동일성을 ULP 한계로 대체하기입니다. PR은 선택지 1을 택했고, 어느 kernel도 고치지 않고 경로를 바꾸는 방식으로 구현했습니다. op 수준 측정에서 WMMA kernel은 128 row 이상에서 측정한 모든 bf16 shape에서 더 느린 쪽이었으므로, 그런 GEMM을 dequantize + hipBLASLt로 보내면 두 경로가 같아지고 실제 prefill도 빨라집니다. 기본 경로를 이전 경로(`MLX_ROCM_WMMA_QMM=1`)와 비교하면 Gemma 3 4B 4-bit의 prefill은 2048 토큰에서 961에서 2815 tok/s로 올라갑니다. + +dispatch 결정은 overlay의 함수 하나, `select_qmm_route`에 모였고, 이 함수가 `QuantizedMatmul::eval_gpu`와 export된 predicate `rocm::quantized_matmul_runs_dequant_gemm` 양쪽에 답을 줍니다. Rust의 적격성 판단은 bridge를 통해 이 predicate에 묻기 때문에, ROCm에서 dense 경로는 `quantized_matmul`이 이미 같은 GEMM을 실행하는 곳에서만 실행됩니다. `814f9efb`에서 실행한 unit의 전체 `make verify-rocm`은 epic #1801 실행 중 처음으로 완전히 통과한 ROCm gate였습니다. + +## 1. 문제 정의 + +### 1.1 실패한 assertion + +``` +panicked at src/lib/mlxcel-core/src/layers.rs:6369:17: +assertion `left == right` failed: dtype 12 bias false: dense GEMM must match qmm bytes +``` + +이슈는 첫 실패 케이스(bf16, bias 없음, x `[1, 1024, 2048]`, 4-bit affine weight `[1024, 2048]`, group 64)를 측정했습니다. 1,048,576개 출력 중 499개(0.048%)가 다릅니다. 465개가 1 ULP, 10개가 2 ULP, 24개가 3에서 35 ULP이며, 모두 크기가 0.01 미만인 출력에서 나왔습니다. 이 범위에서는 상쇄 때문에 bf16 ULP 거리가 커집니다. 가장 큰 절대 오차는 0.25(41.5 대 41.75, 1 ULP)로, 최대 |out| 76.5에 비하면 작습니다. 0에 가까운 원소 하나는 부호가 바뀝니다(-8.5e-6 대 1.6e-5, bit 거리로 28308 ULP). bias 케이스는 실행되지도 않았습니다. 이 실패는 #1806(PR #2030)이 제거한 NVFP4 abort 뒤에 가려져 있었고, 그 이후 모든 전체 ROCm gate에서 실패했습니다. + +### 1.2 Metal에서는 전제가 성립했고 ROCm에서는 성립하지 않은 이유 + +적격성 규칙은 #1994/#2001(PR #2002)에서 왔습니다. affine mode, f16 또는 bf16 입력과 같은 dtype의 scale, 2-D weight, `min_rows` 이상의 row, 32 x 32 출력 tile 512개 초과입니다. 바이트 동일성 주장은 Metal에서 실행한 in-tree sweep에 근거합니다. Metal에서는 두 경로가 같은 반올림으로 dequantize하고 tiling만 다릅니다. ROCm에서는 두 경로가 서로 다른 kernel에 도달합니다. + +- Dense 쪽: `affine_dequantize` 후 `matmul`, 즉 hipBLASLt. +- qmm 쪽, bf16: 장치에 native WMMA가 있고 low-CU iGPU가 아니면 `qmm_wmma_dense_kernel`. 조건은 bf16 x/scales/biases, group 64, 4/6/8 bit, `N % 16 == 0`, `K % 64 == 0`입니다. dequantize는 `affine_dequantize`와 같지만 rocWMMA 16 x 16 x 16 tile을 통해 자체 K 순서로 f32 누적합니다. +- qmm 쪽, f16: WMMA 경로가 없으므로 `affine_dequantize`와 `dequant_rocblas_gemm`, 역시 hipBLASLt입니다. + +`MLX_ROCM_WMMA_QMM=0`으로 테스트가 통과했기 때문에, 원인은 dequantize 차이가 아니라 reduction 순서 차이로 확인되었습니다. 이슈는 #2079에서 고친 CPU stream OpenBLAS 오기록 가능성도 배제했습니다. 테스트의 모든 op는 기본 GPU stream에서 실행됩니다. + +### 1.3 수정 전 production 노출 + +`hardware.rs`의 `prefill_dense_gemm_min_rows_default`는 Apple M1에서만 임계값을 반환합니다. ROCm에서는 운영자가 `MLXCEL_PREFILL_DEQUANT_MIN_M`을 설정해야만 dense 경로가 실행되므로, 이 불일치는 기본 경로 버그가 아니라 opt-in 뒤에 숨은 정확성 구멍이었습니다. + +## 2. 경로 변경을 통한 선택지 1을 고른 이유 + +이슈는 세 선택지 중 하나를 고르라고 요구했고, assertion을 단순히 `cfg`로 끄는 것은 금지했습니다. 선택지 1에는 조건도 붙였습니다. gfx1151에서 bf16 prefill 처리량이 퇴보하지 않아야 하며, op 수준과 실제 모델 양쪽에서 비교하라는 것입니다. + +unit은 선택하기 전에 두 kernel을 측정했습니다. bf16, 4-bit g64, `[1, M, K]` 입력 하나와 `[N, K]` weight, 각 arm 40회 호출을 두 번 번갈아 실행한 평균입니다(`dense`는 매 호출마다 다시 dequantize합니다). + +| M | K x N | qmm ms | dense ms | dense / qmm | differ | +|---:|---|---:|---:|---:|---:| +| 16 | 4096 x 4096 | 0.388 | 0.337 | 0.87 | 79 / 65,536 | +| 16 | 4096 x 14336 | 1.329 | 1.649 | 1.24 | 261 / 229,376 | +| 64 | 4096 x 4096 | 0.413 | 0.488 | 1.18 | 305 / 262,144 | +| 64 | 14336 x 4096 | 1.472 | 1.995 | 1.36 | 1,217 / 262,144 | +| 128 | 4096 x 4096 | 0.805 | 0.570 | 0.71 | 572 / 524,288 | +| 128 | 14336 x 4096 | 2.519 | 2.479 | 0.98 | 2,309 / 524,288 | +| 256 | 4096 x 14336 | 4.022 | 2.274 | 0.57 | 4,153 / 3,670,016 | +| 512 | 4096 x 14336 | 9.187 | 3.769 | 0.41 | 4,242 / 7,340,032 | +| 1024 | 4096 x 1024 | 1.403 | 0.320 | 0.23 | 1,164 / 1,048,576 | +| 1024 | 4096 x 14336 | 15.775 | 4.644 | 0.29 | 12,251 / 14,680,064 | +| 2048 | 14336 x 4096 | 51.491 | 10.046 | 0.20 | 29,524 / 8,388,608 | + +전체 sweep 50개 cell(M 2에서 2048, K x N은 4096 x 4096, 4096 x 1024, 4096 x 14336, 14336 x 4096, 1024 x 3072)에서 WMMA kernel은 64 row 이하의 일부 shape에서만 이겼고, 128 row 이상의 모든 shape에서는 dense가 빨랐습니다(PR 기준 128 row에서 1.0x에서 3.1x, 1024와 2048에서 2.4x에서 5.1x). 모든 cell에서 바이트도 달랐으므로 불일치는 테스트 shape에 한정된 것이 아닙니다. + +정확성과 속도가 같은 방향을 가리켰기 때문에 이 측정으로 선택이 정해졌습니다. + +- **경로 변경을 통한 선택지 1**은 큰 bf16 GEMM을 더 빠르면서 dense 경로와 바이트가 같은 경로로 보냅니다. 변경 후 128 row 이상의 모든 cell에서 `quantized_matmul`과 dense의 차이는 0개이고, 128 row 미만은 바뀌지 않습니다. +- **WMMA kernel을 고치는 선택지 1**(hipBLASLt의 reduction 순서에 맞추기)은 kernel이 지는 바로 그 shape에서 느린 kernel을 유지했을 것이고, kernel을 hipBLASLt 내부 reduction 순서에 묶었을 것입니다. 그 순서는 라이브러리 버전 사이에서 안정된 목표가 아닙니다. +- **선택지 2**(WMMA 장치의 ROCm bf16을 적격성에서 제외)는 테스트를 정직하게 만들었겠지만 bf16 prefill을 느린 kernel에 남겨 두었고, opt-in dense 경로가 ROCm의 bf16 모델에 도움이 될 일도 없었을 것입니다. +- **선택지 3**(ULP 한계)은 dense 경로의 전제가 되는 보장을 포기합니다. 측정된 분포로는 깔끔한 한계를 정하기도 어렵습니다. 2 ULP를 넘는 출력 24개와 bit 거리 28308 ULP의 부호 반전 하나가 있으므로, 0 근처 출력에 맞춘 절대 항이 필요합니다. 이슈는 테스트를 통과시키려고 고른 한계를 명시적으로 경고했습니다. 이 선택지도 bf16 prefill을 느린 kernel에 남겨 둡니다. + +128 row 미만에서는 kernel이 일부 shape에서 여전히 이기므로 그대로 두고, 대신 그곳에서는 dense 경로를 거부합니다(3.3절). 이 부분은 경로 변경이 닿지 않는 곳에만 적용한 선택지 2입니다. + +## 3. 변경 요약 + +`fix/issue-2081-rocm-bf16-dense-gemm`의 세 커밋: + +- **`a73370f2`** `fix(rocm): route large bf16 qmm to hipBLASLt and gate dense prefill`: 경로 함수, ceiling, export된 predicate, bridge 호출, 적격성 변경, ROCm 테스트 블록, 결과 페이지, 문서. +- **`814f9efb`** `fix(rocm): tighten the dense-prefill route check and its test`: 리뷰 후속 작업. predicate가 한 row GEMM(`matmul`이 hipBLASLt가 아니라 gemv로 보냄)을 거부하고, device index를 받아 bridge가 device 0이 아니라 기본 GPU를 확인합니다. route-guard 테스트는 shape마다 예상 적격성을 확인하고, 실제로 다른 shape가 하나 이상 거부되는지 확인하며, 경로와 무관한 assertion을 skip보다 먼저 실행하고, `MLX_NO_HIPBLASLT`를 경로 override로 취급합니다. +- **`c15ff36d`** `docs(rocm): note the dequant cache footprint and more route overrides`: 보안 리뷰 후속 작업. `LOCAL_FIXES.md` 항목 29에 dequantize된 weight LRU의 메모리 점유를 적고, 테스트의 override 목록에 `MLX_ROCM_FORCE_LOW_CU`와 `MLX_ROCM_FORCE_WARP_SIZE`를 추가합니다. + +영역별 파일: + +- Overlay: `patches-rocm/mlx/backend/rocm/quantized/qmm.hip` (`QmmRoute`, `QmmRouteInputs`, `wmma_qmm_env`, `wmma_qmm_max_m`, `select_qmm_route`, `quantized_matmul_runs_dequant_gemm`, dispatch 재배선), `rocm.h` (선언), `no_rocm.cpp` (false를 반환하는 stub), `LOCAL_FIXES.md` 항목 29. +- Bridge: `src/lib/mlxcel-core/cpp/mlx_cxx_bridge.{h,cpp}` (`quantized_matmul_matches_dense_gemm`), `src/lib/mlxcel-core/src/lib.rs` (ffi 선언). +- Rust: `src/lib/mlxcel-core/src/layers.rs` (적격성, doc comment, 테스트), `src/lib/mlxcel-core/src/hardware.rs` (doc comment만). +- 문서: `docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md`, `docs/environment-variables.md`, `docs/mlxcelverse/upstream/README.md` (항목 29는 아직 패키징하지 않음). + +### 3.1 하나의 경로 결정 + +PR 전에는 `QuantizedMatmul::eval_gpu`가 세 개의 분리된 조건으로 경로를 inline에서 결정했습니다. WMMA 블록, dequant GEMM 블록, 그 안의 fp8 블록입니다. PR은 이를 `select_qmm_route(const QmmRouteInputs&, rocm::Device&)`로 옮겼고, 이 함수는 `WmmaDense`, `DequantFp8Gemm`, `DequantGemm`, `Other` 중 하나를 반환합니다. `eval_gpu`는 이제 반환값으로 분기합니다. WMMA shape 조건은 이전과 같은 논리곱입니다. 동작이 바뀌는 곳은 WMMA 분기 안의 조건 하나입니다. + +```cpp +const bool hipblaslt_is_faster = + env != 1 && dequant && in.M >= wmma_qmm_max_m(d) && !fp8(); +``` + +이 조건이 참이면 WMMA kernel에 맞는 GEMM도 `DequantGemm`으로 갑니다. `wmma_qmm_max_m`은 `MLX_ROCM_WMMA_QMM_MAX_M`이 양의 정수면 그 값을, 아니면 `Rdna35` tier에서 128을, 그 밖에서는 `INT_MAX`를 반환합니다. 세 가지 조건이 ceiling이 측정하지 않은 경로로 GEMM을 보내지 못하게 막습니다. + +- `dequant`가 참이어야 합니다. dequant 경로가 사용 가능하고 켜져 있으며(`MLX_ROCM_QMM_DEQUANT_GEMM`이 `0`이 아님) 그 shape에서 선호되거나 강제되어야 합니다. ceiling은 GEMM을 `Other`로 보내지 않습니다. +- `!fp8()`: fallback이 e4m3 경로(fp8을 지원하는 hipBLASLt가 있는 RDNA 4)가 되는 곳에서는 이전처럼 fused kernel이 남습니다. fp8 lambda는 경로를 결정하는 곳에서만 평가되므로, 그 지점에 도달하지 않는 장치는 probe되지 않습니다. +- `env != 1`: `MLX_ROCM_WMMA_QMM=1`은 ceiling을 없앱니다. low-CU iGPU가 아닌 모든 장치에서 이는 이전 dispatch와 같습니다(low-CU iGPU에서는 예전처럼 kernel을 강제로 켭니다). + +### 3.2 Export된 predicate + +`rocm::quantized_matmul_runs_dequant_gemm(device_index, M, N, K, x_dtype, scales_dtype, biases_dtype, group_size, bits)`는 batch 차원이 없는 transposed affine GEMM 하나에 대해 `eval_gpu`가 만드는 `QmmRouteInputs`를 다시 구성합니다. `should_use_dequant_gemm_path`와, qmv가 지원하지 않는 bit 폭에 대한 `force_dequant_gemm` 항도 포함합니다. 그리고 같은 `select_qmm_route`를 호출합니다. `QmmRoute::DequantGemm`이면서 `is_hipblaslt_available()`일 때만 true를 반환합니다. + +- `M < 2`는 거부합니다. `matmul`은 transposed weight에 대한 단일 row를 hipBLASLt가 아니라 gemv로 보냅니다. +- `DequantFp8Gemm`은 거부합니다. fp8 GEMM은 activation과 weight를 e4m3로 반올림하므로 bf16 matmul과 같을 수 없습니다. +- hipBLASLt가 없으면 `dequant_rocblas_gemm`과 `matmul`의 rocBLAS 경로가 서로 다른 rocBLAS 호출로 fallback하므로, hipBLASLt 경우만 주장합니다. + +두 호출자가 같은 함수를 거치므로, 입력을 같은 방식으로 재구성하는 한 predicate는 dispatch와 어긋날 수 없습니다. 환경 변수 override(`MLX_ROCM_WMMA_QMM`, `MLX_ROCM_WMMA_QMM_MAX_M`, `MLX_ROCM_QMM_DEQUANT_GEMM` 등)는 두 답을 함께 움직입니다. + +### 3.3 Bridge와 Rust 적격성 + +bridge의 `quantized_matmul_matches_dense_gemm`은 ROCm bridge 빌드(`MLXCEL_BRIDGE_ROCM_BACKEND`)이고, 런타임 백엔드가 ROCm이며, 기본 장치가 GPU인 경우가 아니면 true를 반환합니다. ROCm에서는 마지막 두 축 앞에 1보다 큰 축이 있는 x를 거부합니다. `QuantizedMatmul`은 그 축들에 대해 batch를 돌리고, batched GEMM은 같은 row에 대해 `matmul`이 실행하는 단일 GEMM이 아니기 때문입니다. 그다음 기본 GPU의 index로 overlay predicate에 묻습니다. `no_rocm.cpp`는 ROCm이 없는 fork 빌드를 위해 predicate를 false로 stub하며, 그런 빌드는 백엔드 확인 때문에 이 지점에 도달하지 않습니다. + +`prefill_dense_gemm_eligible`은 tile 규칙을 유지하고, 더 싼 row와 tile 검사가 먼저 반환한 뒤 마지막 조건으로 bridge 호출을 추가합니다. Metal과 CUDA에서는 호출이 true를 반환하므로 적격성은 이전 규칙과 정확히 같습니다. doc comment에는 acceptance criteria가 요구한 #2081의 근거(불일치 개수, ULP 분포, kernel 경로)가 들어갔고, 이 변경이 강제하는 불변식도 적혀 있습니다. dense 경로는 최적화이므로, 실행되는 곳에서는 `quantized_matmul`이 반환했을 바이트를 반환해야 합니다. + +### 3.4 테스트 + +Metal 시절의 테스트 본문은 의도가 그대로입니다. 경로와 무관한 assertion(정확히 512 tile에서의 narrow-N 거부와 min-rows 거부)이 이제 먼저 실행됩니다. `[1, 1024, 2048]`에서의 바이트 확인은 두 dtype 모두 적격성을 여전히 확인하지만, 측정된 gfx1151 경로가 아닌 ROCm 장치에서 거부된 shape는 실패하지 않고 로그를 남긴 뒤 skip합니다. + +ROCm 전용 helper `prefill_dense_gemm_rocm_route_guard`는 tile 규칙은 통과하지만 다른 경로에 도달하는 shape에서 보장 자체를 확인합니다. + +| dtype | rows | N | gfx1151에서의 경로 | 예상 적격성 | +|---|---:|---:|---|---| +| bf16 | 64 | 8448 | WMMA kernel (ceiling 미만) | 거부 | +| f16 | 64 | 8448 | dequantize + hipBLASLt | 허용 | +| bf16 | 256 | 4096 | dequantize + hipBLASLt (ceiling 이상) | 허용 | + +64 row에서 N 8448은 2 x 264 = 528 tile로 512 tile 하한을 살짝 넘으므로, tile 규칙이 경로 확인을 가리지 않습니다. 모든 ROCm 장치에서 helper는 허용된 shape의 바이트가 같은지 확인합니다. 측정된 경로(`rocm_measured_route()`: gfx1151이고 경로를 움직이는 일곱 개 변수가 모두 unset)에서는 각 예상 적격성과, 실제로 다른 shape가 하나 이상 있었는지도 확인합니다. 그래서 경로 확인을 제거하면 guard가 통과할 수 없습니다. 경로 확인을 우회하면 bf16 64 row에서 실패합니다. 마지막으로 x `[2, 512, 2048]`은 batch 축 때문에 ROCm에서 거부되어야 합니다. + +## 4. 결과 + +### 4.1 모델 prefill + +gfx1151에서 `mlxcel-bench-decode --prompt-tokens {512, 2048} -n 8 --warmup-tokens 4 --ignore-eos`, 바이너리 하나로 `MLX_ROCM_WMMA_QMM=1`("before", 이 장치의 이전 dispatch)과 기본값("after")을 cell마다 ABBA 3회 비교했고, sampler가 다른 GPU 프로세스가 없었음을 확인했습니다. Prefill tok/s, 평균(범위): + +| 모델 | Scales | pp512 before | pp512 after | pp2048 before | pp2048 after | +|---|---|---:|---:|---:|---:| +| gemma-3-4b-it-4bit | bf16 | 1146 (1046-1199) | 2289 (2197-2359) | 961 (948-972) | 2815 (2805-2824) | +| Qwen3-0.6B-4bit | bf16 | 4581 (4230-5024) | 7633 (7147-8012) | 3059 (2894-3219) | 4236 (4027-4566) | +| Qwen3-30B-A3B-4bit | bf16 | 311 (305-315) | 326 (321-334) | 283 (282-284) | 297 (296-299) | +| Meta-Llama-3.1-8B-Instruct-4bit | f16 | 969 (879-1032) | 1008 (994-1026) | 1149 (1143-1153) | 1132 (1124-1137) | + +평균으로 계산하면 Gemma 3 4B는 512 토큰에서 2.0x, 2048에서 2.9x, Qwen3-0.6B는 67%와 38% 빨라집니다. Qwen3-30B-A3B는 두 길이 모두 약 5%인데, attention projection만 이 경로를 타고 expert는 `gather_qmm`을 거치기 때문입니다. + +Llama 3.1 8B는 대조군입니다. scale이 f16이라 어느 GEMM도 바뀐 분기에 도달하지 않으며(f16은 WMMA kernel을 타지 않습니다), 그 변동(512에서 범위가 겹치는 +4.0%, 2048에서 -1.5%)을 설명할 코드 경로가 없습니다. bf16의 개선폭은 이 실행 간 노이즈를 기준으로 읽어야 합니다. + +Decode는 모든 cell에서 변하지 않습니다. M이 1이고, 이는 WMMA kernel을 탄 적이 없습니다. MLX peak 메모리는 Qwen3-0.6B 512 토큰(1.04에서 1.46 GB)을 빼면 0.1 GB 이내입니다. dequantize 경로가 weight matrix마다 bf16 사본을 할당하기 때문이고, 이는 `LOCAL_FIXES.md` 항목 28(PR #2084)이 제한합니다. + +### 4.2 Production 노출: `MLXCEL_PREFILL_DEQUANT_MIN_M` + +ROCm에서 dense 경로는 `MLXCEL_PREFILL_DEQUANT_MIN_M`이 설정되어야만 실행됩니다. 이번 변경 후 그곳에서 허용되는 projection은 `quantized_matmul`이 이미 dequantize + hipBLASLt를 실행하는 것뿐이므로, 변수를 켜도 바이트가 바뀔 수 없습니다. 속도도 의미 있게 바뀌지 않습니다. 두 경로가 같은 GEMM을 실행하고, `quantized_matmul`은 dequantize된 weight cache와 함께 실행합니다. 같은 바이너리로 `MLXCEL_PREFILL_DEQUANT_MIN_M=1024`와 unset을 pp2048에서 ABBA 3회 비교한 결과: gemma-3-4b-it-4bit 2871 대 2852 tok/s, Qwen3-0.6B-4bit 4174 대 4185. + +따라서 production 상황은 다음과 같습니다. + +- **기본 설정(변수 unset)**: ROCm에서 dense 경로는 전후 모두 실행되지 않습니다. 사용자에게 보이는 효과는 bf16 prefill의 더 빠른 `quantized_matmul` 경로입니다. +- **변수 설정**: PR 전에는 운영자가 허용된 projection에서 `quantized_matmul`과 다른 bf16 출력(#2081의 불일치)을 얻을 수 있었습니다. PR 후에는 이 변수가 바이트와 속도 모두에 영향이 없습니다. `docs/environment-variables.md`와 `hardware.rs` doc comment는 ROCm에서 설정해도 이득이 없다고 적고 있고, ROCm 기본값은 꺼진 채로 유지됩니다. + +## 5. 기술적 선택과 그 이유 + +- **세 선택지 중에서 고르기 전에 kernel을 측정했습니다.** 선택지 1에 붙은 처리량 조건이 결정적 근거가 되었습니다. 바이트 동일성을 깬 kernel이 128 row 이상에서는 더 느린 쪽이기도 했습니다. +- **kernel을 다시 쓰지 않고 경로를 바꿨습니다.** 이미 배포된 두 경로 사이에서 shape를 옮기는 것은 새 수치 코드가 필요 없고, 불투명한 라이브러리의 reduction 순서를 맞추는 대신 구조적으로(양쪽이 같은 GEMM) 동일성을 보장합니다. +- **dispatch와 predicate에 함수 하나를 씁니다.** dispatch 조건을 다시 적은 별도 predicate는 누군가 한쪽을 처음 고치는 순간 어긋납니다. `select_qmm_route`를 공유하면 Rust 적격성은 dispatch 자체에 묻는 질문이 됩니다. +- **hipBLASLt 경우만 주장합니다.** predicate는 한 row GEMM, fp8 경로, hipBLASLt 없는 fallback, batched 입력을 거부합니다. 각각은 상위 수준의 경로 이름이 같아도 두 쪽이 다른 코드에 도달하기 때문입니다. +- **ceiling을 측정한 tier로 제한했습니다.** 기본 128 row는 `Rdna35`에만 적용됩니다. RDNA 3, RDNA 4, CDNA는 `MLX_ROCM_WMMA_QMM_MAX_M`을 설정하지 않는 한 모든 row 수에서 kernel을 유지하고, fp8 fallback도 제외되므로, 측정하지 않은 장치가 기본값으로 처리량을 잃지 않습니다. +- **이전 dispatch를 정확히 재현하는 스위치를 남겼습니다.** `MLX_ROCM_WMMA_QMM=1`은 이 장치에서 PR 이전 경로를 복원하며, 그 덕분에 바이너리 하나로 prefill 표의 두 열을 모두 만들 수 있었습니다. +- **ROCm에서 `MLXCEL_PREFILL_DEQUANT_MIN_M` 기본값은 꺼 둡니다.** 경로가 같으므로 ROCm에서 dense 경로가 더할 것이 없다는 것을 4.2절에서 측정했습니다. + +## 6. 검증 + +PR 본문 기준, gfx1151(Radeon 8060S): + +- 대상 테스트가 기본값과 `MLX_ROCM_WMMA_QMM=0`, `MLX_ROCM_WMMA_QMM=1`, `MLX_NO_HIPBLASLT=1`, `MLX_ROCM_FORCE_LOW_CU=1`에서 통과합니다. +- `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` 통과, `cargo test --features rocm --test dead_doc_pointers` 통과. + +Orchestrator 검증: + +- unit이 `814f9efb`에서 전체 `make verify-rocm`을 실행했습니다(gfx1151). 146개 test suite, 11,785개 통과, 0개 실패, 378개 ignored, smoke OK. epic #1801 실행 중 처음으로 완전히 통과한 ROCm gate이며, 직전 baseline 실패가 바로 이 테스트였습니다(PR #2084의 `3c9edea0` gate는 정확히 이 target에서만 실패했습니다). +- head 커밋 `c15ff36d`는 `LOCAL_FIXES.md` 문장 하나와 테스트 override 목록의 환경 변수 이름 두 개만 추가합니다. 그 커밋에서 대상 테스트, clippy, `verify-rocm-overlay`, `verify-fmt`가 통과했습니다. +- Orchestrator는 머지 후 새 빌드에서 최종 gate를 실행합니다. + +## 7. 남은 위험과 검증하지 않은 부분 + +- **경로는 graph를 만들 때 판단합니다.** `prefill_dense_gemm_eligible`은 graph를 만들 때 predicate에 묻고, `dequant_rocblas_gemm`은 eval 시점에 hipBLASLt 사용 가능 여부를 다시 읽습니다. 그 사이에 stream capture가 시작되면 두 쪽이 서로 다른 rocBLAS fallback으로 갈 수 있습니다. 이 백엔드는 HIP graph가 꺼져 있으므로(`use_hip_graphs()`가 false 반환), 이 상황은 decode capture 중에 `MLXCEL_PREFILL_DEQUANT_MIN_M`이 설정되어 있어야 생깁니다. predicate는 hipBLASLt launch 자체가 throw하지 않는다고도 가정합니다. throw하면 각 쪽이 자기 rocBLAS fallback을 탑니다. +- **Dequantize된 weight cache의 메모리 점유.** 이 경로의 LRU(matrix 8개 또는 256 MB, `MLX_ROCM_QMM_DEQUANT_CACHE_SIZE`, `MLX_ROCM_QMM_DEQUANT_CACHE_MAX_BYTES`)는 prefill 동안 모델의 projection 사이에서 hit가 없으므로, 각 projection은 prefill chunk마다 다시 dequantize되고, 마지막 항목들(최대 256 MB)은 prefill 후에도 살아 있습니다. f16 checkpoint는 이미 그렇게 동작했고, 이제 RDNA 3.5의 bf16 checkpoint도 그렇습니다. `MLX_ROCM_QMM_DEQUANT_CACHE_SIZE=0`으로 끌 수 있습니다. 임시 사본은 projection 크기에 비례해 커지며, 호스트에 4B보다 큰 dense bf16 scale checkpoint가 없어서 그런 모델의 peak 메모리는 측정하지 않았습니다. +- **다른 RDNA tier는 측정하지 않았습니다.** ceiling은 `Rdna35` tier 전체(`gfx1150`에서 `gfx1152`)에 적용되지만 gfx1151에서만 측정했습니다. low-CU gfx1152는 강제하지 않는 한 어차피 WMMA kernel을 건너뜁니다. RDNA 3, RDNA 4, CDNA는 이전 dispatch를 유지합니다. gfx1151이 아닌 ROCm 장치에서는 테스트의 바이트 확인이 거부된 shape를 실패 대신 skip하므로 커버리지가 약합니다. +- **64와 128 사이의 row.** 공개된 표에는 그 사이 row 수의 cell이 없습니다. 128은 모든 shape에서 dense가 이긴 첫 측정 row 수입니다. +- **Metal과 CUDA는 이 호스트에서 실행하지 않았습니다.** Rust 적격성과 테스트는 그 경로에서도 바뀌었지만, bridge가 그곳에서 true를 반환하고 ROCm 블록은 ROCm에서만 실행됩니다. +- **Upstream 반영.** 항목 29는 fork용으로 패키징하지 않았습니다. ceiling은 장치 하나에서만 측정했고, export된 predicate는 mlxcel의 dense prefill 확인을 위해 존재합니다. fork PR은 ceiling만 담고, RDNA 3 또는 RDNA 3.5 장치 하나 이상의 측정을 추가해야 합니다. + +## 8. 학습 포인트 + +- **바이트 동일성 주장은 수학이 아니라 kernel에 대한 주장입니다.** #2001의 sweep은 Metal에서 유효했습니다. 다른 백엔드에서는 같은 두 op가 다른 kernel에 도달했고, 전제에는 tile 규칙이 아니라 백엔드 수준의 확인이 필요했습니다. +- **정확성 수정과 성능 문제가 같은 knob을 공유하면 둘 다 먼저 측정해야 합니다.** 여기서는 테스트를 고친 근거가 이슈가 묻지 않았던 2.9x prefill 격차도 없앴습니다. +- **dispatch에 대한 predicate는 dispatch에 직접 묻게 만들어야 합니다.** 조건을 Rust에 옮겨 적지 않고 같은 경로 선택을 실행하는 함수를 export했기 때문에, 환경 변수 override나 이후 tier가 경로를 옮겨도 적격성이 정직하게 유지됩니다. +- **guard 테스트에는 반드시 잡아야 할 mutation을 줘야 합니다.** route-guard는 측정된 경로에서 실제로 다른 shape가 하나 이상 있는지 확인하므로, 경로 확인을 지우면 테스트가 조용히 통과하지 않고 실패합니다. + +Refs: #2081, #1801, #1994, #2001, #2002, #1806, #2030, #2079, #2084. diff --git a/docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md b/docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md new file mode 100644 index 000000000..21c9090ee --- /dev/null +++ b/docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md @@ -0,0 +1,48 @@ +# bf16 quantized matmul route on ROCm: Radeon 8060S (gfx1151), 2026-09-30 + +Why `layers::tests::prefill_dense_gemm_matches_qmm_bytes_where_eligible` failed for bf16 on gfx1151, what was measured to choose the fix, and what the fix did to prefill (issue #2081, part of #1801). The fix is `patches-rocm/LOCAL_FIXES.md` item 29 plus a route check in `prefill_dense_gemm_eligible`. + +## Environment + +Same host as the [decode baseline](rocm-baseline-gfx1151-2026-09-30.md): AMD Ryzen AI MAX+ 395 with Radeon 8060S (`gfx1151`), ROCm 10.0.0 with HIP 7.15.26333, Debian 13. mlxcel at `7278397a` (main) plus this change, MLX pin `81ba1c6a`, overlay NripeshN/mlx `rocm-support` at `75915908` plus `LOCAL_FIXES.md` items 1 to 29. `--features rocm`. Every measured run waited until `/sys/class/kfd/kfd/proc` had been empty for 30 s, and a sampler checked it once per second during the run; one op-level run that another GPU process overlapped was discarded. Compiler activity was not logged. + +## The mismatch + +bf16, no bias, x `[1, 1024, 2048]`, 4-bit g64 affine weight `[1024, 2048]`: 499 of 1,048,576 outputs of `dequantize` + `matmul` differed from `quantized_matmul` (465 by 1 ULP, 10 by 2, 24 by more, all on outputs below 0.01 in magnitude). `quantized_matmul` ran `qmm_wmma_dense_kernel`, which dequantizes with the same rounding as `affine_dequantize` but accumulates through rocWMMA 16 x 16 x 16 tiles in its own K order; the dense side ran hipBLASLt. f16 has no WMMA route, so its `quantized_matmul` also dequantizes and calls hipBLASLt, and it matched. + +## Op level: WMMA kernel against dequantize + hipBLASLt + +bf16, 4-bit g64, one `[1, M, K]` input against an `[N, K]` weight. `qmm` is `quantized_matmul` on main (always the WMMA kernel for these shapes); `dense` is `dequantize` then `matmul`, which re-dequantizes on every call. Mean of 40 calls per arm in two alternated rounds; `differ` counts outputs whose bytes differ. + +| M | K x N | qmm ms | dense ms | dense / qmm | differ | +|---:|---|---:|---:|---:|---:| +| 16 | 4096 x 4096 | 0.388 | 0.337 | 0.87 | 79 / 65,536 | +| 16 | 4096 x 14336 | 1.329 | 1.649 | 1.24 | 261 / 229,376 | +| 64 | 4096 x 4096 | 0.413 | 0.488 | 1.18 | 305 / 262,144 | +| 64 | 14336 x 4096 | 1.472 | 1.995 | 1.36 | 1,217 / 262,144 | +| 128 | 4096 x 4096 | 0.805 | 0.570 | 0.71 | 572 / 524,288 | +| 128 | 14336 x 4096 | 2.519 | 2.479 | 0.98 | 2,309 / 524,288 | +| 256 | 4096 x 14336 | 4.022 | 2.274 | 0.57 | 4,153 / 3,670,016 | +| 512 | 4096 x 14336 | 9.187 | 3.769 | 0.41 | 4,242 / 7,340,032 | +| 1024 | 4096 x 1024 | 1.403 | 0.320 | 0.23 | 1,164 / 1,048,576 | +| 1024 | 4096 x 14336 | 15.775 | 4.644 | 0.29 | 12,251 / 14,680,064 | +| 2048 | 14336 x 4096 | 51.491 | 10.046 | 0.20 | 29,524 / 8,388,608 | + +Across the 50 cells of the sweep (M 2 to 2048; K x N of 4096 x 4096, 4096 x 1024, 4096 x 14336, 14336 x 4096, 1024 x 3072) the WMMA kernel won only on some shapes at 64 rows or fewer, and dense was faster on every shape from 128 rows up. So byte identity and speed pointed the same way: route large bf16 GEMMs to dequantize + hipBLASLt. With the change, `quantized_matmul` differs from dense in 0 outputs of every cell from 128 rows up, and below 128 rows it is unchanged. + +## Model prefill + +`mlxcel-bench-decode --prompt-tokens {512, 2048} -n 8 --warmup-tokens 4 --ignore-eos`, one binary, the default against `MLX_ROCM_WMMA_QMM=1` (which removes the ceiling, the old dispatch on this device), three ABBA runs per cell. Prefill tok/s, mean (range): + +| Model | Scales | pp512 before | pp512 after | pp2048 before | pp2048 after | +|---|---|---:|---:|---:|---:| +| gemma-3-4b-it-4bit | bf16 | 1146 (1046-1199) | 2289 (2197-2359) | 961 (948-972) | 2815 (2805-2824) | +| Qwen3-0.6B-4bit | bf16 | 4581 (4230-5024) | 7633 (7147-8012) | 3059 (2894-3219) | 4236 (4027-4566) | +| Qwen3-30B-A3B-4bit | bf16 | 311 (305-315) | 326 (321-334) | 283 (282-284) | 297 (296-299) | +| Meta-Llama-3.1-8B-Instruct-4bit | f16 | 969 (879-1032) | 1008 (994-1026) | 1149 (1143-1153) | 1132 (1124-1137) | + +Decode is unchanged in every cell (M is 1, which never took the WMMA kernel). MLX peak memory is within 0.1 GB of before except Qwen3-0.6B at 512 tokens, 1.04 to 1.46 GB, because the dequantize route allocates a bf16 copy of each weight matrix, which `LOCAL_FIXES.md` item 28 bounds. Qwen3-30B-A3B gains least because only its attention projections take this path; its experts go through `gather_qmm`. Llama 3.1 8B is the control: its scales are f16, so no GEMM of it reaches the changed code, and its -1.5% at 2048 tokens has no path to explain it. No dense bf16-scale checkpoint above 4B was available on the host, so peak memory for one is not measured here; the dequantized-weight cache holds 8 matrices or 256 MB, so every model measured here dequantizes each projection again per prefill chunk, and the transient copies grow with the projection size. + +## The dense prefill path on ROCm + +`MLXCEL_PREFILL_DEQUANT_MIN_M=1024` against unset, same binary, pp2048, three ABBA runs: gemma-3-4b-it-4bit 2871 against 2852 tok/s, Qwen3-0.6B-4bit 4174 against 4185. With the route change the dense path is eligible on ROCm only where `quantized_matmul` already runs the same GEMM, so turning it on changes neither bytes nor speed there. diff --git a/docs/environment-variables.md b/docs/environment-variables.md index c8eefbfca..655d1673f 100644 --- a/docs/environment-variables.md +++ b/docs/environment-variables.md @@ -117,7 +117,7 @@ Both server entry points implement llama-server b10621's Vertex AI custom-contai | `MLX_MAX_MB_PER_BUFFER` (MLX-native, Metal and CUDA) | positive integer, nominally MB | Metal: unset, so MLX's own table applies (40 on base/Pro/phone, 50 on Max and Ultra); mlxcel raises the budget for decode steps only, see `MLXCEL_DECODE_MB_PER_BUFFER`. CUDA: `1000` on compute capability 12.1 (GB10) for the family whose every measured workload gains from it (`model_type` `laguna`, #1798), untouched for every other checkpoint and elsewhere (MLX's own table: 400 on A100, 1000 on H100, B200 and consumer Blackwell, 25 on 12.1, 100 for anything unlisted) | Input-size budget of one captured CUDA graph. The counter is not bytes: MLX sums `array::data_size()` over the graph's input arrays, an element count, and compares `count >> 20` against this value, so on the 25 default any op reading an array over 26.2M elements commits its own graph. That is every 256-expert NVFP4 expert stack on Laguna (about 120 `gather_qmm` per token) and every 4-bit `lm_head` or tied embedding over 26.2M packed words (Qwen 3.5 4B, Llama 3.1 8B, Gemma 3 4B). On a model where that happens per layer, capture at 25 costs more than graphs off; whether raising the budget helps is a per-family measurement (see `MLX_MAX_OPS_PER_BUFFER` above). On Metal the same element-count budget decides when a command buffer is committed. Setting it here pins it for prefill and decode alike and turns off the decode-only switch below. An explicit value always overrides. | | `MLXCEL_DECODE_MB_PER_BUFFER` | `0`/`off`/`false`/`no` (disable), positive integer | `1000` on pre-M5 Apple Silicon (the same gate as `MLX_MAX_OPS_PER_BUFFER`), off elsewhere | Metal command-buffer input budget applied during decode steps only. MLX commits a command buffer once the element count of its distinct inputs, shifted right by 20, passes `MLX_MAX_MB_PER_BUFFER` (40-50 by default). With `MLX_MAX_OPS_PER_BUFFER` raised to 1000 that budget is the only cap that binds, and in decode, where every token reads the whole weight set, the default commits a buffer every one to two layers (about 23 per token on command-r7b 4-bit) and idles the GPU at each boundary. mlxcel raises the budget around pipelined decode only (the generate loops and the server's lookahead decode) and leaves prefill and synchronous decode steps on the device default, through a runtime override in the `mlx/backend/metal/device.cpp` overlay. A synchronous step encodes and then waits, so one large buffer there only removes the overlap between CPU encoding and GPU execution. Measured on M1 Ultra (500-token prompt, 128 generated tokens, three interleaved runs per cell), decode at 1000 versus the default: command-r7b 4-bit +7%, Llama 3.1 8B 4-bit +5.6%, Qwen2.5 7B 4-bit +8%, Gemma 3n E4B +3%, Granite 4.0 H Tiny +10%, Qwen3-30B-A3B +20%, Mixtral 8x7B +21%, Llama 3.1 8B bf16 +17%, Gemma 3 4B flat. Applying the same value to prefill as well would roughly double peak memory on long prompts (2048 tokens: Qwen2.5 7B 4-bit 6.0 to 12.6 GB, Qwen3-30B-A3B 19.8 to 36.2 GB) and cost up to 2.7% prefill throughput, which is why it is scoped to decode. Speculative draft/verify loops are not covered yet. Ignored when `MLX_MAX_MB_PER_BUFFER` is set. Measurement: `docs/benchmark_results/metal-mb-per-buffer-m1ultra-2026-09-21.md`. | | `MLXCEL_MAMBA1_SCAN_KERNEL` | `0` disables; anything else or unset keeps the default | on (Metal; CUDA for Jamba) | Fused Metal kernel for the Mamba1 selective scan in Jamba (#2005) and Mamba / Falcon-Mamba (#2007). One kernel walks every timestep of a layer with the state in float32 registers (one simdgroup per channel, one lane per state element), for prefill and decode alike, replacing a Rust loop of small ops per timestep. Measured on M1 Ultra, Jamba reasoning 3B 4-bit: prefill +45% at 512 tokens and +54% at 1024 and 2048 (about 830 to 1330 tok/s), decode about +10%, prefill peak memory 4.11 to 3.15 GB at 2048 tokens. The float32 state makes results differ from the graph scan, which rounds the bf16 state every step: teacher-forced perplexity improves slightly (15.357 to 15.271 on the WikiText-2 excerpt) and top-1 choices differ only where the graph path's top two were within one logit. Set to `0` to force the graph scan (A/B, rollback). Checked on every call. On CUDA (#1981) Jamba uses a port that rounds every step in the activation dtype exactly as the graph scan does, so its output is bit-identical to the graph scan when all scan inputs share one dtype (otherwise the graph scan runs); on GB10 it cut a 3.3k-token Jamba chat request from 4.3 s to 1.5 s. Mamba / Falcon-Mamba still use the graph scan on CUDA, and ROCm always does. Measurements: `docs/benchmark_results/jamba-mamba1-scan-kernel-m1ultra-2026-09-28.md`, `docs/benchmark_results/jamba-mamba1-scan-kernel-gb10-2026-09-30.md`. | -| `MLXCEL_PREFILL_DEQUANT_MIN_M` | `0`/`off`/`false`/`no` (disable), positive integer row count | `1024` on M1-generation Apple Silicon, off on every other Apple generation, CUDA and ROCm | Input row count (every axis but the last, so a server batch counts all its rows) at which an affine 4-bit projection whose scales share the input's dtype (f16 or bf16) runs as `dequantize` + dense matmul instead of `quantized_matmul`. Only projections whose output has more than 512 tiles of 32 x 32 qualify (at 1024 rows, an output wider than 512): there the two return identical bytes, while narrower outputs tile differently and stay on `quantized_matmul` (#2001). Measured on M1 Ultra, prefill at 1024 rows: Llama 3.1 8B +9.7 to +10.0%, command-r7b +7.7%, Phi-3 mini +10.2%, Qwen2.5 7B +0.6 to +1.5%, Gemma 2 2B +0.6%, Mixtral 8x7B +0.6%; bf16-scale Qwen3 1.7B +15.1%, Qwen3-30B-A3B +5.5%, Gemma 4 12B +3.0%, Gemma 4 E4B +0.9%, Gemma 3 4B +0.5%. At 2048 rows every one of them gains, +0.8 to +18.0%. Below 1024 several lose (at 512 rows Qwen2.5 7B -4.2%, Gemma 4 E4B -2.7%, Gemma 2 2B -1.3%), which sets the default. Costs 0.1 to 0.6 GB of prefill peak memory. The server's default 512-token prefill chunks stay below the default; the CLI's 2048-token chunks cross it. Other generations are unmeasured, so the default is off there; set a row count to opt in. Measurement: `docs/benchmark_results/prefill-dense-gemm-m1ultra-2026-09-27.md`. | +| `MLXCEL_PREFILL_DEQUANT_MIN_M` | `0`/`off`/`false`/`no` (disable), positive integer row count | `1024` on M1-generation Apple Silicon, off on every other Apple generation, CUDA and ROCm | Input row count (every axis but the last, so a server batch counts all its rows) at which an affine 4-bit projection whose scales share the input's dtype (f16 or bf16) runs as `dequantize` + dense matmul instead of `quantized_matmul`. Only projections whose output has more than 512 tiles of 32 x 32 qualify (at 1024 rows, an output wider than 512): there the two return identical bytes, while narrower outputs tile differently and stay on `quantized_matmul` (#2001). On ROCm a projection also qualifies only where `quantized_matmul` itself runs dequantize + hipBLASLt, the one route that returns the same bytes (#2081); there the dense path repeats that GEMM without the backend's dequantized-weight cache, so setting this variable on ROCm gains nothing. Measured on M1 Ultra, prefill at 1024 rows: Llama 3.1 8B +9.7 to +10.0%, command-r7b +7.7%, Phi-3 mini +10.2%, Qwen2.5 7B +0.6 to +1.5%, Gemma 2 2B +0.6%, Mixtral 8x7B +0.6%; bf16-scale Qwen3 1.7B +15.1%, Qwen3-30B-A3B +5.5%, Gemma 4 12B +3.0%, Gemma 4 E4B +0.9%, Gemma 3 4B +0.5%. At 2048 rows every one of them gains, +0.8 to +18.0%. Below 1024 several lose (at 512 rows Qwen2.5 7B -4.2%, Gemma 4 E4B -2.7%, Gemma 2 2B -1.3%), which sets the default. Costs 0.1 to 0.6 GB of prefill peak memory. The server's default 512-token prefill chunks stay below the default; the CLI's 2048-token chunks cross it. Other generations are unmeasured, so the default is off there; set a row count to opt in. Measurement: `docs/benchmark_results/prefill-dense-gemm-m1ultra-2026-09-27.md`. | | `MLXCEL_HEADROOM_FACTOR` | positive `f64` | `1.20` | Runtime/activation headroom multiplier used by the unified memory estimator (`mlxcel inspect`, `--estimate-memory`, `--recommend-quant`). Positive values `<= 1.0` disable the headroom term; invalid or non-positive values warn and fall back to the default. Override only for calibration runs — see the in-code recipe in `src/execution/memory_estimate.rs`. | | `MLXCEL_CACHE_DIR` | directory path | `$HOME/.cache/mlxcel` | Root for mlxcel's on-disk caches. The tokenizer language-analysis disk cache (language-bias features) lives under `tokenizer-scripts/`, and the location-independent global model store lives under `models//` when `MLXCEL_MODELS_DIR` and the store-root flag (`--model-store-root` on the servers, `--models-dir` on the subcommands) are both unset. | | `MLXCEL_MODELS_DIR` | directory path | unset (falls back to `${MLXCEL_CACHE_DIR:-$HOME/.cache/mlxcel}/models`) | Dedicated model-store root. Snapshots live directly at `$MLXCEL_MODELS_DIR//` with no `models/` subdir, so the whole store can sit on a separate volume without dragging the tokenizer-script cache along. Read by `mlxcel download`, the `-m/--model` resolver (`generate` / `serve` / `inspect` / `run`), the `mlxcel-server -m/--model` resolver, and `list` / `rm`. Resolution precedence for the models root: the CLI flag (`--model-store-root ` on `mlxcel-server` / `mlxcel serve` since #1438 reserved `--models-dir` for b10621 router mode; still `--models-dir ` on the `download` / `list` / `rm` / `generate` subcommands), then `MLXCEL_MODELS_DIR`, then `${MLXCEL_CACHE_DIR:-$HOME/.cache/mlxcel}/models`. (`download --local-dir ` is separate: it writes the snapshot verbatim at that exact path.) | @@ -215,6 +215,20 @@ the server accepts. The stream stays wedged behind the kernel either way, so a watchdog failure still means restarting the process. It does not apply under `MLX_EVENT_BLOCKING`, whose blocking waits return promptly on a fault but have no poll loop to time. +`MLX_ROCM_WMMA_QMM` picks how `quantized_matmul` runs a bf16 affine GEMM with +more than one row (4-, 6- or 8-bit, group 64): unset uses the fused WMMA kernel +where the device has native WMMA and is not a low-CU iGPU, `0` never uses it, +and `1` uses it wherever the shape fits. `MLX_ROCM_WMMA_QMM_MAX_M` is the row +count from which such a GEMM instead goes to dequantize + hipBLASLt, where +that route is enabled (`MLX_ROCM_QMM_DEQUANT_GEMM` not `0`) and is not the fp8 +path; it defaults to `128` on RDNA 3.5 (`gfx1150` to `gfx1152`) and to no +ceiling elsewhere, a value that is not a positive integer keeps that default, +and `MLX_ROCM_WMMA_QMM=1` ignores it. On gfx1151 dequantize + +hipBLASLt was faster than the WMMA kernel on every bf16 shape measured from 128 +rows up, which raised bf16 prefill there by 5% (Qwen3-30B-A3B, whose experts do +not take this path) to 2.9x (Gemma 3 4B at 2048 tokens); see +`docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md` +(lablup/mlxcel#2081). ## OpenXLA / StableHLO backend variables diff --git a/docs/mlxcelverse/upstream/README.md b/docs/mlxcelverse/upstream/README.md index cf8ab7dfd..f90fae2a1 100644 --- a/docs/mlxcelverse/upstream/README.md +++ b/docs/mlxcelverse/upstream/README.md @@ -34,6 +34,7 @@ Each directory holds the patch (`git format-patch` output, authored by the maint | 7 (GPU fault reporting) | Deferred. It is built on `Event::error()`, which ml-explore/mlx added in #3742 (`06f154bc`, 2026-08-17); the fork's MLX base predates it, so the patch does not apply. The wait changes that end a hang on a faulted queue (`HipEvent::wait`, `AtomicEvent`, `CommandEncoder::synchronize`) could be split out for the current fork, but that would be a different change from the one mlxcel ships and verified. Package it once the fork merges MLX at or past #3742. | | 9 (expert-batched gather qmv opt-in) | Held. It disables a path that returns wrong bf16 results on gfx1151 but has no root cause yet; LOCAL_FIXES.md says to propose it only after one. | | 15 (`SearchSorted`) | Not applicable yet: the primitive arrived in ml-explore/mlx#4035 (`5ec30acd`, 2026-08-10), after the fork's MLX base, so the fork has no `SearchSorted` to implement. Package it once the fork merges MLX at or past #4035. | +| 29 (bf16 GEMMs leave the WMMA kernel from 128 rows on RDNA 3.5) | Not packaged yet. The row ceiling was measured on gfx1151 only, and the exported `quantized_matmul_runs_dequant_gemm` exists for mlxcel's dense prefill check; a fork PR would carry the ceiling alone, with measurements from at least one more RDNA 3 or RDNA 3.5 device. | ## What was verified, and what the submitter still has to do diff --git a/src/lib/mlx-cpp/patches-rocm/LOCAL_FIXES.md b/src/lib/mlx-cpp/patches-rocm/LOCAL_FIXES.md index 70e7cb6f8..017fb0363 100644 --- a/src/lib/mlx-cpp/patches-rocm/LOCAL_FIXES.md +++ b/src/lib/mlx-cpp/patches-rocm/LOCAL_FIXES.md @@ -37,7 +37,8 @@ The fork branched from upstream on 2026-06-26 (`39886de4`). Moving its ROCm code 24. **`SliceUpdate` and `DynamicSliceUpdate` no longer donate a source that is still referenced.** Both `SliceUpdate::eval_gpu` (`mlx/backend/rocm/indexing.hip`) and `DynamicSliceUpdate::eval_gpu` (`mlx/backend/gpu/primitives.cpp`) gave the source's buffer to the output whenever the buffer had one owner (`in.data_shared_ptr().use_count() == 1`), without checking whether the source array itself was still referenced. A second handle to the same array, or an unevaluated node that takes it as an input (a `Slice`, a `Copy`), holds the array but not its buffer, so the update was written into an array that could still be read. Upstream's `array::is_donatable()` requires both counts to be one, and Metal and CUDA reach it through `copy_gpu`, whose `set_copy_output_data` donates only when `is_donatable(in, out)` holds. Both primitives now call `copy_gpu` exactly as upstream CUDA's `SliceUpdate::eval_gpu` and upstream's `DynamicSliceUpdate::eval_gpu` do. `DynamicSliceUpdate` also forced donation whenever `rocm::graph_active()` was set, for KV cache accumulation under HIP graph capture; that override is removed, because it donated even a buffer other arrays shared, and it could not fire: `graph_active()` is set only by the `CommandEncoder` constructor when `use_hip_graphs()` is true, `use_hip_graphs()` returns a constant `false` in `mlx/backend/rocm/device.cpp`, and `MLX_GRAPH_PREFILL_REPLAY` is read only behind `use_hip_graphs()`. If HIP graphs are turned back on, accumulation under capture has to be solved in the capture, not by writing into a referenced source. It was a live bug in mlxcel, measured on gfx1151: DeepSeek-V4's `PoolingCache::accumulate_windows` returns windows that lazily read the previous remainder rows while the same call writes the new remainder over those rows, and the model's cache barrier evaluates the buffer first, so a chunked prefill that completes a window across a remainder and leaves a new one read the new tail instead (`tiny_model_chunked_prefill_with_pool_remainder_matches_cpu` moved the first logit by about 9e-3 relative against the CPU; 1e-5 after the fix); and a prompt-cache `ModelStateSnapshot` of a `RotatingKVCache`, a lazy `copy` of the live keys, took the next token's ring overwrite when a one-token suffix wrapped the ring (`rotating_steady_state_snapshot_survives_the_next_wrap_write`; the existing `rotating_inplace_warmup_and_wrap_keep_shared_snapshots_intact` also fails on the old build with `MLXCEL_KV_INPLACE_WRITE=0`). `tests/rocm_slice_update_source.rs` holds a source, runs a GPU update through `slice_update` (None), `slice_update_reduce` (Sum, Max) and the test-only `slice_update_dynamic` bridge, and checks the source is unchanged; every held-source case fails on the old build. A source dropped before evaluation is still donated, which is the KV cache's own update pattern (`self.keys = slice_update(self.keys, ...)`), and the single-token FP16 decode write goes through `inplace_slice_write` rather than `SliceUpdate`. Decode throughput on gfx1151 (`scripts/bench_decode.sh`, pp512/tg128, interleaved runs with no other GPU process or compiler running): Qwen3-0.6B-4bit 279.0 and 279.8 tok/s before, 279.5 and 279.9 after; Meta-Llama-3.1-8B-Instruct-4bit 35.68, 35.67 and 35.82 before, 35.90 and 35.69 after; no change beyond the run-to-run spread (under 1%). Applies to the fork; to be proposed there (upstreaming candidate, lablup/mlxcel#1813). 25. **`MLX_ROCM_FFT_CACHE_SIZE` validated.** The fork's `LRUBytesKeyCache` constructor in `mlx/backend/rocm/lru_cache.h` read its capacity variable with an unchecked `std::stoul`, and its one env-driven user is the hipFFT plan cache in `fft.hip` (entry 18). Per value, confirmed with a host g++ 14 probe of the header: `0` was accepted, and the first `put()` then ran `while (size() >= capacity_)` on an empty list and called `back()` and `pop_back()` on it, which is undefined behavior (the probe threw `std::bad_alloc`; with `-D_GLIBCXX_ASSERTIONS` it aborts on `!this->empty()`); `abc`, an empty value and an overflowing value such as `99999999999999999999999` threw `std::invalid_argument` or `std::out_of_range` whose message is only `stoul`, out of the static's initializer in `FFT::eval_gpu`, so the static stayed uninitialized and every later FFT threw again; `-1` was taken silently as `SIZE_MAX` (a cache that never evicts), `8abc` silently as 8, and `0x10` as 0, with the same undefined behavior as `0` (on gfx1151 the first GPU rfft failed with `std::bad_alloc` for both). `capacity_from_env()` in the same header now parses it the way `gpu_watchdog_seconds()` in `device.cpp` parses `MLX_ROCM_GPU_WATCHDOG_SECS`: `std::strtol` base 10, the whole string must be the number, and the value must be 1 to `INT_MAX` (the bound CUDA's `int` read implies). Unset gives the default silently; anything else prints one stderr line, `[ROCm] ignoring invalid MLX_ROCM_FFT_CACHE_SIZE="" (expected a positive integer); using the default 128`, and uses the default, once, because the static then initializes. Both `LRUBytesKeyCache` and `LRUCache(size_t)` now throw `std::runtime_error("LRUCache requires capacity > 0.")` for a capacity of 0, as upstream CUDA's `LRUCache` does (`mlx/backend/cuda/lru_cache.h` at the pin). Upstream CUDA reads `MLX_CUDA_FFT_CACHE_SIZE` with `atoi`, so there `0` or junk throws that error at the first FFT and `-1` becomes `SIZE_MAX`; that was not copied. The comment on `fft_plan_cache()` in `fft.hip` states the rule; editing that file was also what made a warm build recompile it, because at the time each HIP object depended only on its own `.hip` source (fixed in item 26). `tests/rocm_fft_cache_env.rs` runs one child process per value (`0`, `-1`, `abc`, empty, `8abc`, `0x10`, an overflow past `long`, `2147483648`, `16`, `1`, `2147483647`, unset) and checks the warning count and an rfft against the CPU stream. Applies to the fork; to be proposed there (upstreaming candidate, lablup/mlxcel#1813). 27. **CPU-stream BLAS runs single-threaded over fine-grained memory.** On an APU the fork's allocator (`mlx/backend/rocm/allocator.cpp`, `unified_malloc`) gives every array fine-grained device memory, CPU-stream arrays included, and `Buffer::raw_ptr()` hands the CPU that same pointer. MLX's CPU backend writes BLAS and LAPACK results straight into it (`cblas_sgemm` in `mlx/backend/cpu/gemms/cblas.cpp`, also `conv.cpp`, `masked_mm.cpp` and the LAPACK primitives), and multithreaded OpenBLAS returned wrong output columns there, differently from call to call. Measured on the gfx1151 host (Ryzen AI MAX+ 395, Debian OpenBLAS 0.3.29 pthread, 32 threads): a standalone C `cblas_sgemm` of `[1, 2880] x [2880, 2880]^T` was wrong in 37 of 50 calls with inputs and output from `hipExtMallocWithFlags(hipDeviceMallocFinegrained)`, in 10 of 300 with only the output there, and exact in every call with malloc or `hipHostMalloc` buffers or with `OPENBLAS_NUM_THREADS` at 1, 2 or 4. Through MLX the wrong columns were each the last column of one OpenBLAS thread's share of the output (shares of 93 columns: 650, 1022, 1859, ...), the column whose 64-byte line the next thread also writes. This was lablup/mlxcel#2072: `tests/rocm_mxfp4_quant.rs` failed `qmm 2880x2880 M=1` in 3 of 8 runs here (7 of 10 in the issue), and the fault was in its CPU reference, not the GPU kernel; with the tensors dumped from a failing run, the GPU result matched an independent dequantize-and-dot of the packed bytes and the CPU `matmul` did not. `unified_malloc` now calls `openblas_set_num_threads(1)` once, on its first fine-grained allocation, which precedes any CPU-stream BLAS call writing such a buffer; the symbol is declared weak, so a build against another BLAS links and changes nothing. The cost is CPU-stream BLAS speed: the CPU reads fine-grained memory slowly, and that gemv went from about 0.2 s to about 1.4 s per call. GPU work and quantized or half-precision CPU matmuls (MLX's own kernels, no BLAS) are unaffected, so model inference under `MLXCEL_DEVICE=cpu` barely touches BLAS. Keeping BLAS multithreaded by giving it a cacheable scratch output would need a change at every BLAS call site in MLX's CPU backend and was not taken. `tests/rocm_cpu_blas_finegrained.rs` repeats that gemv on the CPU stream against an f64 host reference; with this change reverted it failed in 5 of 5 runs, and with it `quantized_matmul_matches_dequantized_reference` passed 50 runs in a row. Applies to the fork; to be proposed there (upstreaming candidate, lablup/mlxcel#1813). -28. **In-flight command batches bounded by what they allocate, and the cache limit enforced.** Measured on the gfx1151 host at pp512/tg128 (lablup/mlxcel#2062), the MLX peak (live buffers, `get_peak_memory`) was 20.60 GB for Meta-Llama-3.1-8B-Instruct-4bit, whose weights take 4.75 GB once loaded, and 23.56 GB for Qwen3-30B-A3B-4bit (17.17 GB), and the device-wide `mem_info_vram_used` rose by 21.19 GB and 24.09 GB. The excess was not buffer cache (0.4 to 1.7 GB at every phase boundary) but live transients inside the measured prefill: every buffer an operation allocates is kept by `CommandEncoder::add_temporary` until its batch's completion handler runs, the eager path (`use_hip_graphs()` is off) committed only every `MLX_MAX_OPS_PER_BUFFER` ops (2000) and never registered those commits with MLX's scheduler, so neither the scheduler's task cap nor `set_memory_limit` ever made the host wait, and the host could encode far ahead of the GPU, holding the transients of many batches at once. For the 8B, an f16 checkpoint, most of it is the f16 weight copy `QuantizedMatmul`'s dequantize-and-GEMM path allocates per matrix (with `MLX_ROCM_QMM_DEQUANT_GEMM=0` the peak was 6.78 GB, at a twentieth of the prefill speed); the MoE model is bf16, whose affine 4-bit matmuls take the fused WMMA kernel and allocate no copy (the variable left its peak at 23.56 GB), so its excess is other operations' outputs held the same way, which the bound below removes too. `MLX_MAX_OPS_PER_BUFFER=50` alone changed nothing, because more commits do not stop the host from running ahead. Now `gpu::eval` (`mlx/backend/rocm/eval.cpp`) adds what each primitive allocated during `eval_gpu` to its encoder's open batch, read from a thread-local byte counter in the allocator (`rocm::thread_allocated_bytes()`, so frees on the worker thread do not disturb it), and `CommandEncoder::maybe_commit` (`mlx/backend/rocm/device.cpp`) commits once the open batch has allocated a quarter of `MLX_ROCM_MAX_INFLIGHT_MB` (default 1024); after that commit and after the commit `gpu::finalize` makes at the end of every eval and `async_eval`, it records a HIP event, and blocks on the oldest such batch while the committed ones still exceed the budget. The wait is `hipEventSynchronize` (the device is in blocking-sync mode, so the thread sleeps), or a polled wait that gives up at `MLX_ROCM_GPU_WATCHDOG_SECS` when that is set; a faulted stream ends either with the fault, which is recorded as the stream's error. Nothing is tracked while a decode-step capture or any stream capture is in progress. The events are raw `hipEvent_t` owned by the encoder, not `HipEvent`: `HipEvent` returns its handle to a function-local static pool when destroyed, an encoder can be destroyed after that pool at process exit, and a first version built on it aborted `tests/rocm_mxfp4_quant.rs` at exit with `malloc_consolidate(): unaligned fastbin chunk detected`. `0` restores the old behavior. `tests/rocm_inflight_bound.rs` runs 40 f16 4-bit 8192x8192 matmuls over 512 rows in one evaluation (5 GiB of dequantized weight copies): the peak rose by 1.75 GiB with the default and by 6.25 GiB with `MLX_ROCM_MAX_INFLIGHT_MB=0`, which fails its 3 GiB bound. With the default the peaks were 6.14 GB and 18.57 to 18.58 GB over three harness runs each, decode throughput was unchanged (median 37.57 against 37.02 tok/s and 61.69 against 61.79 tok/s) and the 8B's prefill lost 1.8% (the MoE model's prefill varies more than that from run to run); budgets from 256 to 4096 MiB traded peak for prefill speed within a few percent (the table is in `docs/benchmark_results/rocm-memory-gfx1151-2026-09-30.md`). Separately, `RocmAllocator::set_cache_limit` stored `max_pool_size_` and nothing read it, so MLX's `set_cache_limit` (mlxcel's `MLXCEL_CACHE_LIMIT`) bounded nothing, while the exact-size cache (`min_utilization` 1.0) keeps every buffer size it has seen until `clear_cache()`. `malloc_async` now trims the cache on a cache miss once it exceeds `max_pool_size_`, the only place the footprint grows, so active plus cache stays within the live set plus the limit plus the one request; it trims to three quarters of the limit, so one round of blocking `hipFree` calls covers several misses of a workload whose shapes keep changing; `free()` and cache hits still never call `hipFree`, which the fork avoids for the training reasons in its comment there. The limit still defaults to the memory limit (76.8 GiB here), so the trim changes nothing until someone sets a limit; mlxcel sets 2 GiB on ROCm builds. `memory::tests::cache_limit_bounds_the_free_buffer_cache` in mlxcel-core frees 64 buffers of distinct sizes under a 1 MiB limit and checks the cache stays under 8 MiB; with the trim removed it held 49.8 MB and failed. The count is of allocations, not only of transients (a batch that grows the KV cache counts too), and a batch's buffers are dropped by the worker thread just after its event completes, so the bound is approximate. Applies to the fork; to be proposed there (upstreaming candidate, lablup/mlxcel#1813). +28. **In-flight command batches bounded by what they allocate, and the cache limit enforced.** Measured on the gfx1151 host at pp512/tg128 (lablup/mlxcel#2062), the MLX peak (live buffers, `get_peak_memory`) was 20.60 GB for Meta-Llama-3.1-8B-Instruct-4bit, whose weights take 4.75 GB once loaded, and 23.56 GB for Qwen3-30B-A3B-4bit (17.17 GB), and the device-wide `mem_info_vram_used` rose by 21.19 GB and 24.09 GB. The excess was not buffer cache (0.4 to 1.7 GB at every phase boundary) but live transients inside the measured prefill: every buffer an operation allocates is kept by `CommandEncoder::add_temporary` until its batch's completion handler runs, the eager path (`use_hip_graphs()` is off) committed only every `MLX_MAX_OPS_PER_BUFFER` ops (2000) and never registered those commits with MLX's scheduler, so neither the scheduler's task cap nor `set_memory_limit` ever made the host wait, and the host could encode far ahead of the GPU, holding the transients of many batches at once. For the 8B, an f16 checkpoint, most of it is the f16 weight copy `QuantizedMatmul`'s dequantize-and-GEMM path allocates per matrix (with `MLX_ROCM_QMM_DEQUANT_GEMM=0` the peak was 6.78 GB, at a twentieth of the prefill speed); the MoE model is bf16, whose affine 4-bit matmuls took the fused WMMA kernel at the time and allocated no copy (item 29 moved its prefill projections to the dequantize path) (the variable left its peak at 23.56 GB), so its excess is other operations' outputs held the same way, which the bound below removes too. `MLX_MAX_OPS_PER_BUFFER=50` alone changed nothing, because more commits do not stop the host from running ahead. Now `gpu::eval` (`mlx/backend/rocm/eval.cpp`) adds what each primitive allocated during `eval_gpu` to its encoder's open batch, read from a thread-local byte counter in the allocator (`rocm::thread_allocated_bytes()`, so frees on the worker thread do not disturb it), and `CommandEncoder::maybe_commit` (`mlx/backend/rocm/device.cpp`) commits once the open batch has allocated a quarter of `MLX_ROCM_MAX_INFLIGHT_MB` (default 1024); after that commit and after the commit `gpu::finalize` makes at the end of every eval and `async_eval`, it records a HIP event, and blocks on the oldest such batch while the committed ones still exceed the budget. The wait is `hipEventSynchronize` (the device is in blocking-sync mode, so the thread sleeps), or a polled wait that gives up at `MLX_ROCM_GPU_WATCHDOG_SECS` when that is set; a faulted stream ends either with the fault, which is recorded as the stream's error. Nothing is tracked while a decode-step capture or any stream capture is in progress. The events are raw `hipEvent_t` owned by the encoder, not `HipEvent`: `HipEvent` returns its handle to a function-local static pool when destroyed, an encoder can be destroyed after that pool at process exit, and a first version built on it aborted `tests/rocm_mxfp4_quant.rs` at exit with `malloc_consolidate(): unaligned fastbin chunk detected`. `0` restores the old behavior. `tests/rocm_inflight_bound.rs` runs 40 f16 4-bit 8192x8192 matmuls over 512 rows in one evaluation (5 GiB of dequantized weight copies): the peak rose by 1.75 GiB with the default and by 6.25 GiB with `MLX_ROCM_MAX_INFLIGHT_MB=0`, which fails its 3 GiB bound. With the default the peaks were 6.14 GB and 18.57 to 18.58 GB over three harness runs each, decode throughput was unchanged (median 37.57 against 37.02 tok/s and 61.69 against 61.79 tok/s) and the 8B's prefill lost 1.8% (the MoE model's prefill varies more than that from run to run); budgets from 256 to 4096 MiB traded peak for prefill speed within a few percent (the table is in `docs/benchmark_results/rocm-memory-gfx1151-2026-09-30.md`). Separately, `RocmAllocator::set_cache_limit` stored `max_pool_size_` and nothing read it, so MLX's `set_cache_limit` (mlxcel's `MLXCEL_CACHE_LIMIT`) bounded nothing, while the exact-size cache (`min_utilization` 1.0) keeps every buffer size it has seen until `clear_cache()`. `malloc_async` now trims the cache on a cache miss once it exceeds `max_pool_size_`, the only place the footprint grows, so active plus cache stays within the live set plus the limit plus the one request; it trims to three quarters of the limit, so one round of blocking `hipFree` calls covers several misses of a workload whose shapes keep changing; `free()` and cache hits still never call `hipFree`, which the fork avoids for the training reasons in its comment there. The limit still defaults to the memory limit (76.8 GiB here), so the trim changes nothing until someone sets a limit; mlxcel sets 2 GiB on ROCm builds. `memory::tests::cache_limit_bounds_the_free_buffer_cache` in mlxcel-core frees 64 buffers of distinct sizes under a 1 MiB limit and checks the cache stays under 8 MiB; with the trim removed it held 49.8 MB and failed. The count is of allocations, not only of transients (a batch that grows the KV cache counts too), and a batch's buffers are dropped by the worker thread just after its event completes, so the bound is approximate. Applies to the fork; to be proposed there (upstreaming candidate, lablup/mlxcel#1813). +29. **bf16 GEMMs of 128 rows or more leave the fused WMMA kernel on RDNA 3.5.** `QuantizedMatmul::eval_gpu` (`mlx/backend/rocm/quantized/qmm.hip`) sent every single bf16 affine GEMM with more than one row (4-, 6- or 8-bit, group 64, `N % 16 == 0`, `K % 64 == 0`) to `qmm_wmma_dense_kernel` on a device with native WMMA that is not a low-CU iGPU, however many rows it had. On gfx1151 that kernel lost to the dequantize + hipBLASLt path it bypasses on every bf16 shape measured from 128 rows up (1.0x to 3.1x at 128 rows, 1.6x to 2.4x at 256, 2.4x to 5.1x at 1024 and 2048; it won only on some shapes at 2 to 64 rows), so bf16 prefill ran at a fraction of what the same device did for f16 checkpoints (lablup/mlxcel#2081). It also made `quantized_matmul` and `dequantize` + `matmul` disagree for bf16 (499 of 1,048,576 outputs at `[1024, 2048] x [1024, 2048]`, a reduction-order difference), which is what mlxcel's dense prefill path checks. The dispatch decision is now one function, `select_qmm_route`, used by `eval_gpu`; it keeps the kernel below `MLX_ROCM_WMMA_QMM_MAX_M` rows (default 128 on the RDNA 3.5 tier, no ceiling on the unmeasured tiers) and sends larger GEMMs to the bf16 dequantize + hipBLASLt route, except where that route would be the fp8 path (RDNA 4), which keeps the kernel as before. `MLX_ROCM_WMMA_QMM=1` removes the ceiling, which is the old dispatch on any device that is not a low-CU iGPU (there it forces the kernel on, as it always did). `rocm.h` exports `quantized_matmul_runs_dequant_gemm` (stubbed in `no_rocm.cpp`), which answers from the same function whether a single GEMM of a given shape takes dequantize + hipBLASLt, for mlxcel to decide when its dense prefill path returns the same bytes. Same binary, `MLX_ROCM_WMMA_QMM=1` against the default, prefill tok/s over three ABBA runs: Gemma 3 4B 4-bit 1146 to 2289 at 512 tokens and 961 to 2815 at 2048, Qwen3-0.6B 4-bit 4581 to 7633 and 3059 to 4236, Qwen3-30B-A3B 4-bit (MoE, only the attention projections take this path) 311 to 326 and 283 to 297; decode unchanged, and MLX peak memory within 0.1 GB except Qwen3-0.6B at 512 tokens (1.04 to 1.46 GB), since the route allocates a bf16 weight copy per matrix that item 28 bounds. The route's dequantized-weight LRU (8 matrices or 256 MB, `MLX_ROCM_QMM_DEQUANT_CACHE_SIZE`, `MLX_ROCM_QMM_DEQUANT_CACHE_MAX_BYTES`) gets no hits across a model's projections in prefill but still keeps its last entries alive afterwards, as it already did for f16 checkpoints; `MLX_ROCM_QMM_DEQUANT_CACHE_SIZE=0` turns it off. `docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md` has the tables. Applies to the fork; the ceiling was measured on gfx1151 only. ## Fixes to the fork's build diff --git a/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/no_rocm.cpp b/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/no_rocm.cpp index 896bf2719..5f2f1d813 100644 --- a/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/no_rocm.cpp +++ b/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/no_rocm.cpp @@ -50,6 +50,19 @@ std::vector moe_swiglu_sorted_vjp( throw std::runtime_error("moe_swiglu_sorted_vjp requires ROCm"); } +bool quantized_matmul_runs_dequant_gemm( + int, + int, + int, + int, + Dtype, + Dtype, + std::optional, + int, + int) { + return false; +} + } // namespace rocm namespace fast { diff --git a/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/quantized/qmm.hip b/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/quantized/qmm.hip index 161efbe98..61fa692bd 100644 --- a/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/quantized/qmm.hip +++ b/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/quantized/qmm.hip @@ -35,6 +35,7 @@ #include #include #include +#include #include #include #include @@ -822,6 +823,116 @@ inline bool should_use_dequant_gemm_path( return false; } +// Where QuantizedMatmul::eval_gpu sends one GEMM. Shared by the dispatch and by +// rocm::quantized_matmul_runs_dequant_gemm, so the answer that predicate gives +// can never drift from what eval_gpu launches (lablup/mlxcel#2081). +enum class QmmRoute { + WmmaDense, // qmm_wmma_dense_kernel, fused dequant + rocWMMA + DequantFp8Gemm, // dequant_fp8_gemm (weight and activation in e4m3) + DequantGemm, // dequantize (cached) + dequant_rocblas_gemm in the x dtype + Other, // qmv, batched qmv and the tiled kernels below +}; + +struct QmmRouteInputs { + int M; + int N; + int K; + int batch_count; + int x_batch_count; + int w_batch_count; + bool transpose; + QuantizationMode mode; + Dtype x_dtype; + Dtype scales_dtype; + std::optional biases_dtype; + int group_size; + int bits; + bool force_dequant_gemm; + bool should_prefer_dequant; +}; + +// MLX_ROCM_WMMA_QMM: unset = default, "0" = off, "1" = force on. +inline int wmma_qmm_env() { + static const int value = [] { + const char* e = std::getenv("MLX_ROCM_WMMA_QMM"); + if (!e) + return -1; + if (e[0] == '0') + return 0; + if (e[0] == '1') + return 1; + return -1; + }(); + return value; +} + +// Rows at and above which a GEMM the dispatch would otherwise hand to +// dequantize + hipBLASLt stays there instead of taking qmm_wmma_dense_kernel +// (lablup/mlxcel#2081). Measured on gfx1151 (Radeon 8060S, RDNA 3.5), bf16, +// 4-bit g64: the fused kernel wins only on some shapes up to 64 rows, while +// dequantize + hipBLASLt is faster on every shape tried from 128 rows up +// (1.0x to 3.1x at 128 to 512 rows, 2.4x to 5.1x at 1024 and 2048), which +// took Gemma 3 4B 4-bit prefill from 961 to 2815 tok/s at 2048 tokens. Other +// tiers were not measured and keep the kernel at every row count unless +// MLX_ROCM_WMMA_QMM_MAX_M sets a ceiling; MLX_ROCM_WMMA_QMM=1 removes it. +inline int wmma_qmm_max_m(rocm::Device& d) { + static const int env_value = + parse_positive_int_env("MLX_ROCM_WMMA_QMM_MAX_M", -1); + if (env_value > 0) { + return env_value; + } + return detect_rocm_hw_info(d).tier == rocm::RocmArchTier::Rdna35 + ? 128 + : std::numeric_limits::max(); +} + +inline QmmRoute select_qmm_route(const QmmRouteInputs& in, rocm::Device& d) { + const bool single = + in.batch_count == 1 && in.x_batch_count == 1 && in.w_batch_count == 1; + const bool affine = in.mode == QuantizationMode::Affine; + const bool dequant = affine && d.is_rocblas_available() && + use_rocblas_dequant_path() && + (in.force_dequant_gemm || in.should_prefer_dequant); + // fp8 e4m3 path (RDNA4 prefill), capability-gated on the hipBLASLt probe. + // Evaluated only where it decides the route, so a device that never reaches + // it is not probed. + auto fp8 = [&] { + return dequant && in.x_dtype == bfloat16 && single && in.M >= 64 && + rocm::device_has_fp8_gemm(d.hip_device()); + }; + + // Low-CU iGPUs (gfx1152 4-8 CU) skip WMMA-QMM unless forced: fat 64x128 + // tiles thrash the tiny L2 and have produced garbage/NaN under APU memory + // pressure. + const int env = wmma_qmm_env(); + const bool want_wmma = + env == 1 || (env < 0 && !detect_rocm_hw_info(d).is_low_cu_igpu); + const bool wmma_shape = in.transpose && affine && + in.x_dtype == bfloat16 && in.scales_dtype == bfloat16 && + (!in.biases_dtype.has_value() || *in.biases_dtype == bfloat16) && + single && in.M > 1 && (in.N % 16 == 0) && in.group_size == 64 && + (in.K % 64 == 0) && + (in.bits == 4 || in.bits == 6 || in.bits == 8) && d.has_native_wmma(); + if (want_wmma && wmma_shape) { + // The ceiling only hands the GEMM to the bf16 dequant + hipBLASLt path it + // was measured against; where the fallback would be the fp8 path the + // fused kernel stays. + const bool hipblaslt_is_faster = + env != 1 && dequant && in.M >= wmma_qmm_max_m(d) && !fp8(); + if (!hipblaslt_is_faster) { + return QmmRoute::WmmaDense; + } + return QmmRoute::DequantGemm; + } + if (fp8()) { + return QmmRoute::DequantFp8Gemm; + } + if (dequant) { + return QmmRoute::DequantGemm; + } + return QmmRoute::Other; +} + struct DequantCacheKey { std::uintptr_t w_ptr; std::uintptr_t scales_ptr; @@ -3280,6 +3391,59 @@ __global__ void __launch_bounds__(128) qmm_wmma_dense_kernel( } // namespace rocm +namespace rocm { + +bool quantized_matmul_runs_dequant_gemm( + int device_index, + int M, + int N, + int K, + Dtype x_dtype, + Dtype scales_dtype, + std::optional biases_dtype, + int group_size, + int bits) { + // One row is never claimed: `matmul` sends a single row with a transposed + // weight to gemv (gemms/gemv.hip), not hipBLASLt. + if (M < 2 || N < 1 || K < 1 || + !(x_dtype == float16 || x_dtype == bfloat16)) { + return false; + } + auto& d = rocm::device( + mlx::core::Device(mlx::core::Device::gpu, device_index)); + // The same inputs QuantizedMatmul::eval_gpu derives for one transposed + // affine GEMM with no batch dimensions. + const bool bits_supported_by_qmv = + bits == 2 || bits == 4 || bits == 8 || bits == 5 || bits == 6; + const bool should_prefer_dequant = + should_use_dequant_gemm_path(M, N, K, 1, true, true, false, d); + const QmmRoute route = select_qmm_route( + QmmRouteInputs{ + M, + N, + K, + 1, + 1, + 1, + true, + QuantizationMode::Affine, + x_dtype, + scales_dtype, + biases_dtype, + group_size, + bits, + !bits_supported_by_qmv, + should_prefer_dequant}, + d); + // dequant_rocblas_gemm and matmul's gemm_rocblas both try hipBLASLt first + // and fall back to different rocBLAS calls, so only the hipBLASLt case is + // claimed. The claim assumes the hipBLASLt launch itself does not throw, + // which would send each side to its own rocBLAS fallback. + return route == QmmRoute::DequantGemm && is_hipblaslt_available(); +} + +} // namespace rocm + void QuantizedMatmul::eval_gpu(const std::vector& inputs, array& out) { auto& s = stream(); auto& d = rocm::device(s.device); @@ -3330,37 +3494,34 @@ void QuantizedMatmul::eval_gpu(const std::vector& inputs, array& out) { bool force_dequant_gemm = !transpose_ || !bits_supported_by_qmv || ((batch_count > 1) && !can_use_batched_qmv) || (w.ndim() > 2 && !w_singleton_batch && !can_use_batched_qmv); - bool dequant_gemm_supported_mode = (mode_ == QuantizationMode::Affine); bool should_prefer_dequant = should_use_dequant_gemm_path( M, N, K, batch_count, non_batched, transpose_, can_use_batched_qmv, d); + const QmmRoute route = select_qmm_route( + QmmRouteInputs{ + M, + N, + K, + batch_count, + x_batch_count, + w_batch_count, + transpose_, + mode_, + x.dtype(), + scales.dtype(), + biases.has_value() ? std::optional(biases->dtype()) + : std::nullopt, + group_size_, + bits_, + force_dequant_gemm, + should_prefer_dequant}, + d); // Fused WMMA quantized GEMM (q4/q6/q8): graph-addable single node, no dequant - // materialization. Default ON where the device has native WMMA (opt out with - // MLX_ROCM_WMMA_QMM=0); dispatch below falls to rocBLAS when WMMA is absent or - // the shape doesn't qualify. Single (non-batched) bf16 affine prefill. - // Low-CU iGPUs (gfx1152 4–8 CU): skip WMMA-QMM — fat 64×128 tiles thrash the - // tiny L2 and have produced garbage/NaN under APU memory pressure; QMV/dequant - // paths are safer. Force on with MLX_ROCM_WMMA_QMM=1. + // materialization. Single (non-batched) bf16 affine GEMM on a device with + // native WMMA, below the row ceiling in select_qmm_route (opt out with + // MLX_ROCM_WMMA_QMM=0, force on wherever it fits with MLX_ROCM_WMMA_QMM=1). { - static const int wmma_qmm_env = [] { - const char* e = std::getenv("MLX_ROCM_WMMA_QMM"); - if (!e) - return -1; // default - if (e[0] == '0') - return 0; - if (e[0] == '1') - return 1; - return -1; - }(); - auto hw_early = detect_rocm_hw_info(d); - const bool want_wmma = (wmma_qmm_env == 1) || - (wmma_qmm_env < 0 && !hw_early.is_low_cu_igpu); - if (want_wmma && transpose_ && mode_ == QuantizationMode::Affine && - x.dtype() == bfloat16 && scales.dtype() == bfloat16 && - (!biases.has_value() || biases->dtype() == bfloat16) && - batch_count == 1 && x_batch_count == 1 && w_batch_count == 1 && M > 1 && - (N % 16 == 0) && group_size_ == 64 && (K % 64 == 0) && - (bits_ == 4 || bits_ == 6 || bits_ == 8) && d.has_native_wmma()) { + if (route == QmmRoute::WmmaDense) { dim3 block(128, 1, 1); dim3 grid((M + 63) / 64, (N + 127) / 128, 1); const bool hb = biases.has_value(); @@ -3391,9 +3552,7 @@ void QuantizedMatmul::eval_gpu(const std::vector& inputs, array& out) { // Dequant + rocBLAS GEMM path // Disable with MLX_ROCM_QMM_DEQUANT_GEMM=0 if needed - if (dequant_gemm_supported_mode && d.is_rocblas_available() && - use_rocblas_dequant_path() && - (force_dequant_gemm || should_prefer_dequant)) { + if (route == QmmRoute::DequantGemm || route == QmmRoute::DequantFp8Gemm) { if (!((x_batch_count == 1) || (x_batch_count == batch_count))) { throw std::runtime_error( "Unsupported x batch shape for dequant GEMM fallback"); @@ -3409,9 +3568,7 @@ void QuantizedMatmul::eval_gpu(const std::vector& inputs, array& out) { // fp8 e4m3 path (RDNA4 prefill): dequantize the weight straight to e4m3 and // cast the activation, then run the GEMM on fp8 matrix cores. Capability- // gated — devices without e4m3 kernels stay on the bf16 dequant path below. - if ((mode_ == QuantizationMode::Affine) && (x.dtype() == bfloat16) && - (batch_count == 1) && (x_batch_count == 1) && (w_batch_count == 1) && - (M >= 64) && rocm::device_has_fp8_gemm(d.hip_device())) { + if (route == QmmRoute::DequantFp8Gemm) { dequant_fp8_gemm( enc, transpose_, diff --git a/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/rocm.h b/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/rocm.h index c67e1c1f8..56f937967 100644 --- a/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/rocm.h +++ b/src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/rocm.h @@ -8,6 +8,7 @@ #include "mlx/utils.h" #include +#include #include namespace mlx::core::rocm { @@ -51,4 +52,23 @@ MLX_API std::vector moe_swiglu_sorted_vjp( const array& dy, StreamOrDevice s = {}); +// True when QuantizedMatmul on GPU `device_index`, for a transposed affine +// weight and one [M, K] input with no batch dimensions (M >= 2), dequantizes +// the weight in the input dtype and runs the GEMM through hipBLASLt. That is +// the path `matmul` takes for a dense f16 or bf16 weight, so on this route +// `dequantize` + `matmul` returns the same bytes as `quantized_matmul`; on the +// others (the fused WMMA kernel, the fp8 path, qmv) it does not. Reads the same +// route selection QuantizedMatmul uses, including its environment overrides +// (lablup/mlxcel#2081). +MLX_API bool quantized_matmul_runs_dequant_gemm( + int device_index, + int M, + int N, + int K, + Dtype x_dtype, + Dtype scales_dtype, + std::optional biases_dtype, + int group_size, + int bits); + } // namespace mlx::core::rocm diff --git a/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp b/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp index bdd267f3b..15300be13 100644 --- a/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp +++ b/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp @@ -5,6 +5,10 @@ #include "../../mlx-cpp/turbo/gpu_backend.h" #include "mlx/primitives.h" +#ifdef MLXCEL_BRIDGE_ROCM_BACKEND +// `quantized_matmul_runs_dequant_gemm` (lablup/mlxcel#2081). +#include "mlx/backend/rocm/rocm.h" +#endif #include "sampling.h" // Gumbel-max categorical sampling kernel (#900). #include "sampling_rejection.h" // Dual-pivot rejection sampling kernel (#901). @@ -5202,6 +5206,52 @@ bool bitlinear_kernel_available() { return mlxcel::gpu_kernel_backend() != mlxcel::GpuKernelBackend::None; } +// See the header. Only ROCm routes `quantized_matmul` to kernels whose bytes +// depend on the shape in a way the tile rule on the Rust side does not +// capture, so every other backend answers true and keeps its eligibility. +bool quantized_matmul_matches_dense_gemm( + const MlxArray& x, + const MlxArray& weight, + const MlxArray& scales, + const MlxArray* biases, + int32_t group_size, + int32_t bits +) { +#ifdef MLXCEL_BRIDGE_ROCM_BACKEND + using namespace mlx::core; + const Device device = default_device(); + if (mlxcel::gpu_kernel_backend() != mlxcel::GpuKernelBackend::Rocm || + device.type != Device::gpu) { + return true; + } + const auto& xs = x.inner.shape(); + const auto& ws = weight.inner.shape(); + if (xs.size() < 2 || ws.size() != 2) { + return false; + } + // QuantizedMatmul batches over every axis before the last two, and a + // batched GEMM is not the single GEMM `matmul` runs on the same rows. + for (size_t i = 0; i + 2 < xs.size(); ++i) { + if (xs[i] != 1) { + return false; + } + } + std::optional biases_dtype = + biases ? std::optional(biases->inner.dtype()) : std::nullopt; + return rocm::quantized_matmul_runs_dequant_gemm( + device.index, xs[xs.size() - 2], ws[0], xs.back(), x.inner.dtype(), + scales.inner.dtype(), biases_dtype, group_size, bits); +#else + (void)x; + (void)weight; + (void)scales; + (void)biases; + (void)group_size; + (void)bits; + return true; +#endif +} + // Top-p (nucleus) filtering. // Not compiled: MLX v0.31.x Scan primitive (cumsum) lacks output_shapes, // which causes "CumSum cannot infer output shapes" when used inside diff --git a/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.h b/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.h index 81a654ee6..292977790 100644 --- a/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.h +++ b/src/lib/mlxcel-core/cpp/mlx_cxx_bridge.h @@ -1420,6 +1420,22 @@ std::unique_ptr rocm_fault_probe_array(int32_t kind); // (issues #1803, #1862). bool bitlinear_kernel_available(); +// Whether `quantized_matmul(x, weight, scales, biases)` with a transposed +// affine weight runs the same dense GEMM that `dequantize` + `matmul` runs, so +// the two return the same bytes (lablup/mlxcel#2081). On ROCm this reads the +// QuantizedMatmul route (the fused WMMA kernel, the fp8 path and qmv differ; +// dequantize + hipBLASLt matches), and x with a batch axis above 1 answers +// false. Every other backend, and a ROCm build running on the CPU, answers +// true: their eligibility rule is the output-tile count on the Rust side. +bool quantized_matmul_matches_dense_gemm( + const MlxArray& x, + const MlxArray& weight, + const MlxArray& scales, + const MlxArray* biases, + int32_t group_size, + int32_t bits +); + // Fused sampling: top-k + top-p + min-p on the untempered distribution, then // one temperature scaling, then the categorical draw, in a single function // call to minimize FFI round-trips (chain order per issue #1379). diff --git a/src/lib/mlxcel-core/src/hardware.rs b/src/lib/mlxcel-core/src/hardware.rs index 0e34886cf..e4176df92 100644 --- a/src/lib/mlxcel-core/src/hardware.rs +++ b/src/lib/mlxcel-core/src/hardware.rs @@ -644,7 +644,9 @@ pub const PREFILL_DENSE_GEMM_ENV: &str = "MLXCEL_PREFILL_DEQUANT_MIN_M"; /// -2.5% / -0.2%, with +0.5 to 0.8 GB of peak memory, so M5 stays off. M2 /// through M4 differ in GPU microarchitecture and were not measured, and later /// or non-Apple devices report `Unknown`; they keep `quantized_matmul` until a -/// measurement adds them here. +/// measurement adds them here. On ROCm there is nothing to add: the dense path +/// is only eligible where `quantized_matmul` already runs the same dequantize + +/// hipBLASLt GEMM (#2081), which it does with a cached dequantized weight. #[must_use] pub fn prefill_dense_gemm_min_rows_default(r#gen: AppleSiliconGen) -> Option { (r#gen == AppleSiliconGen::M1).then_some(PREFILL_DENSE_GEMM_MIN_ROWS) diff --git a/src/lib/mlxcel-core/src/layers.rs b/src/lib/mlxcel-core/src/layers.rs index 4eb2340bc..2acfd53ea 100644 --- a/src/lib/mlxcel-core/src/layers.rs +++ b/src/lib/mlxcel-core/src/layers.rs @@ -693,28 +693,44 @@ fn fused_qk_norm_enabled_from(value: Option<&str>) -> bool { } } -// ── Fused residual-add + RMSNorm (issue #905) ──────────────────────────────── +// ── Dense-GEMM prefill (issues #1994, #2001, #2081) ────────────────────────── -/// Default for the fused residual-add + RMSNorm decode path. -/// -/// **This is the one place to flip if the measurement does not justify the -/// fusion.** `MLXCEL_FUSED_ADD_RMSNORM=0` disables it at runtime without a -/// rebuild; setting this constant to `false` makes off the default and -/// `MLXCEL_FUSED_ADD_RMSNORM=1` the opt-in. -/// /// Whether [`UnifiedLinear`] should run this quantized projection as -/// `dequantize` + dense matmul (issues #1994, #2001): an affine weight whose -/// scales share the input's dtype (f16 or bf16), a 2-D weight, at least +/// `dequantize` + dense matmul (issues #1994, #2001, #2081): an affine weight +/// whose scales share the input's dtype (f16 or bf16), a 2-D weight, at least /// `min_rows` input rows (the product of every axis but the last, so a server -/// batch counts all of its rows), and more than 512 output tiles of 32 x 32. +/// batch counts all of its rows), more than 512 output tiles of 32 x 32, and a +/// backend whose `quantized_matmul` runs the same GEMM for this shape +/// ([`ffi::quantized_matmul_matches_dense_gemm`]). The dense path is an +/// optimization, so wherever it runs it must return the bytes +/// `quantized_matmul` would have. /// /// With matching dtypes the dequantized weight equals what `quantized_matmul` -/// reconstructs in registers, so the two differ only if they tile the output -/// differently. Above 512 tiles they do not: an in-tree sweep over f16 and bf16, -/// M 1024 and 2048, K 2048 to 4096 and N 256 to 4096 found identical bytes in -/// every such cell, and differences only at or below 512 tiles (for example -/// M 1024 with N 512 or narrower), where `qmm_splitk` targets about 512 -/// threadgroups. Those narrow projections stay on `quantized_matmul`. +/// reconstructs in registers, so on Metal the two differ only if they tile the +/// output differently. Above 512 tiles they do not: an in-tree sweep over f16 +/// and bf16, M 1024 and 2048, K 2048 to 4096 and N 256 to 4096 found identical +/// bytes in every such cell, and differences only at or below 512 tiles (for +/// example M 1024 with N 512 or narrower), where `qmm_splitk` targets about 512 +/// threadgroups. Those narrow projections stay on `quantized_matmul`. On Metal +/// and CUDA the backend check always passes, so the tile count is the rule. +/// +/// On ROCm the tile count says nothing: which kernel `quantized_matmul` runs +/// depends on dtype, rows and device, and only one of its routes matches the +/// dense path. Measured on gfx1151 (#2081), bf16, x `[1, 1024, 2048]` against +/// a 4-bit g64 `[1024, 2048]` weight: the fused `qmm_wmma_dense_kernel` +/// accumulates through rocWMMA 16 x 16 x 16 tiles in its own K order, and +/// 499 of 1,048,576 outputs differed from `dequantize` + hipBLASLt (465 by +/// 1 ULP, 10 by 2, 24 by more, all on outputs below 0.01 in magnitude where +/// cancellation inflates the ULP distance). f16 matched because its +/// `quantized_matmul` also dequantizes and calls hipBLASLt. The ROCm overlay +/// now hands a bf16 GEMM of 128 rows or more on RDNA 3.5 to that same +/// dequantize + hipBLASLt route, which was faster there on every shape +/// measured (1.0x to 3.1x at 128 rows, up to 5.1x at 2048; LOCAL_FIXES +/// item 29), so bf16 matches at those shapes too. The shapes that +/// still take the WMMA kernel, the fp8 path or qmv (fewer rows, other devices, +/// the `MLX_ROCM_*` overrides, a batch axis above 1) fail the backend check and +/// stay on `quantized_matmul`. See +/// `docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md`. pub(crate) fn prefill_dense_gemm_eligible( x: &MlxArray, weight: &QuantizedWeight, @@ -735,7 +751,22 @@ pub(crate) fn prefill_dense_gemm_eligible( .map(|&d| d as i64) .product(); let tiles = ((rows + 31) / 32) * ((w_shape[0] as i64 + 31) / 32); - rows >= min_rows && tiles > DENSE_GEMM_MIN_OUTPUT_TILES + if rows < min_rows || tiles <= DENSE_GEMM_MIN_OUTPUT_TILES { + return false; + } + // SAFETY: every reference is a live array owned by `x` or `weight`, and + // `biases_ptr` is either null or points into `weight`, which outlives the + // call; the bridge only reads shapes and dtypes. + unsafe { + ffi::quantized_matmul_matches_dense_gemm( + x, + &weight.weight, + &weight.scales, + weight.biases_ptr(), + weight.group_size, + weight.bits, + ) + } } /// Output-tile count (32 x 32 tiles of the `[rows, N]` result) at or below @@ -780,6 +811,15 @@ fn dense_gemm( } } +// ── Fused residual-add + RMSNorm (issue #905) ──────────────────────────────── + +/// Default for the fused residual-add + RMSNorm decode path. +/// +/// **This is the one place to flip if the measurement does not justify the +/// fusion.** `MLXCEL_FUSED_ADD_RMSNORM=0` disables it at runtime without a +/// rebuild; setting this constant to `false` makes off the default and +/// `MLXCEL_FUSED_ADD_RMSNORM=1` the opt-in. +/// /// Default-OFF: measured, and the measurement did not justify wiring it on. /// /// Op-level microbench on Apple M1 Ultra (Metal, f16, hidden {2048, 4096, @@ -6337,6 +6377,11 @@ mod tests { /// more than 512 output tiles, the dense path returns the same bytes as /// `quantized_matmul`, with and without a linear bias. Narrow outputs /// (512 tiles or fewer), mixed dtypes and short inputs are not eligible. + /// + /// On ROCm (#2081) the same holds only where `quantized_matmul` runs the + /// dequantize + hipBLASLt route, so the ROCm block below checks the + /// guarantee itself on shapes that reach other routes: any shape the rule + /// accepts must match, and a shape whose bytes differ must be refused. #[test] fn prefill_dense_gemm_matches_qmm_bytes_where_eligible() { let (rows, k) = (1024, 2048); @@ -6344,11 +6389,31 @@ mod tests { let x = ffi::astype(&seeded(&[1, rows, k], 11), dt); let wide = quantized_4bit(&ffi::astype(&seeded(&[1024, k], 7), dt)); let bias = ffi::astype(&seeded(&[1024], 13), dt); - assert!(prefill_dense_gemm_eligible(&x, &wide, 1024), "dtype {dt}"); + // 1024 rows x 512 columns is exactly 512 tiles: MLX tiles the two + // paths differently there, so it must stay on qmm. + let narrow = quantized_4bit(&ffi::astype(&seeded(&[512, k], 7), dt)); + assert!( + !prefill_dense_gemm_eligible(&x, &narrow, 1), + "dtype {dt}: narrow N" + ); assert!( !prefill_dense_gemm_eligible(&x, &wide, 1025), "dtype {dt}: min rows" ); + + // On ROCm eligibility also depends on the device's qmm route + // (#2081). On the measured route both dtypes take dequantize + + // hipBLASLt at 1024 rows; elsewhere a refused shape skips the byte + // check instead of failing it. + let eligible = prefill_dense_gemm_eligible(&x, &wide, 1024); + let rocm = crate::hardware::gpu_backend_kind() == crate::hardware::GpuBackendKind::Rocm; + if rocm && !rocm_measured_route() && !eligible { + eprintln!( + "skipping dtype {dt} byte check: this ROCm device's qmm route differs from dequantize + hipBLASLt (lablup/mlxcel#2081)" + ); + continue; + } + assert!(eligible, "dtype {dt}"); for b in [None, Some(&bias)] { let bias_ptr = b .map(|b| b.as_ref().unwrap() as *const MlxArray) @@ -6373,14 +6438,10 @@ mod tests { b.is_some() ); } + } - // 1024 rows x 512 columns is exactly 512 tiles: MLX tiles the two - // paths differently there, so it must stay on qmm. - let narrow = quantized_4bit(&ffi::astype(&seeded(&[512, k], 7), dt)); - assert!( - !prefill_dense_gemm_eligible(&x, &narrow, 1), - "dtype {dt}: narrow N" - ); + if crate::hardware::gpu_backend_kind() == crate::hardware::GpuBackendKind::Rocm { + prefill_dense_gemm_rocm_route_guard(k); } let x16 = ffi::astype(&seeded(&[1, rows, k], 11), crate::dtype::FLOAT16); @@ -6392,6 +6453,89 @@ mod tests { ); } + /// True on gfx1151, the device the ROCm qmm routes were measured on + /// (#2081), when no environment variable that moves a route is set. + fn rocm_measured_route() -> bool { + crate::rocm_arch::device_gfx_target() == Some("gfx1151") + && [ + "MLX_ROCM_WMMA_QMM", + "MLX_ROCM_WMMA_QMM_MAX_M", + "MLX_ROCM_QMM_DEQUANT_GEMM", + "MLX_ROCM_QMM_DEQUANT_M_THRESHOLD", + "MLX_NO_HIPBLASLT", + "MLX_ROCM_FORCE_LOW_CU", + "MLX_ROCM_FORCE_WARP_SIZE", + ] + .iter() + .all(|name| std::env::var_os(name).is_none()) + } + + /// ROCm half of [`prefill_dense_gemm_matches_qmm_bytes_where_eligible`]: + /// shapes that pass the tile rule but reach different `quantized_matmul` + /// routes. Everywhere, a shape the rule accepts must match. On the + /// measured route, bf16 at 64 rows takes the fused WMMA kernel, whose + /// bytes differ, so the rule must refuse it, while f16 at 64 rows and bf16 + /// at 256 rows take dequantize + hipBLASLt and must be accepted. A batch + /// axis above 1 makes `quantized_matmul` batch the GEMM, so it is refused + /// whatever the bytes. + fn prefill_dense_gemm_rocm_route_guard(k: i32) { + let measured = rocm_measured_route(); + let mut refused_differing = 0; + for (dt, rows, n, expect_eligible) in [ + (crate::dtype::BFLOAT16, 64, 8448, false), + (crate::dtype::FLOAT16, 64, 8448, true), + (crate::dtype::BFLOAT16, 256, 4096, true), + ] { + let x = ffi::astype(&seeded(&[1, rows, k], 17), dt); + let w = quantized_4bit(&ffi::astype(&seeded(&[n, k], 19), dt)); + let want = unsafe { + ffi::quantized_linear_forward( + &x, + &w.weight, + &w.scales, + w.biases_ptr(), + std::ptr::null(), + w.group_size, + w.bits, + &w.mode, + ) + }; + let same = raw_bytes(&dense_gemm(&x, &w, None)) == raw_bytes(&want); + let eligible = prefill_dense_gemm_eligible(&x, &w, 1); + eprintln!( + "rocm route guard: dtype {dt} rows {rows} N {n}: eligible {eligible}, same bytes {same}" + ); + assert!( + !eligible || same, + "dtype {dt} rows {rows} N {n}: eligible, but dense GEMM differs from qmm" + ); + if measured { + assert_eq!( + eligible, expect_eligible, + "dtype {dt} rows {rows} N {n}: eligibility on gfx1151" + ); + } + if !same { + refused_differing += 1; + } + } + if measured { + // The bf16 64-row case must really differ, or this guard would + // pass with the route check removed. + assert!( + refused_differing >= 1, + "no differing shape exercised the route check" + ); + } + + let batched_x = ffi::astype(&seeded(&[2, 512, k], 23), crate::dtype::BFLOAT16); + let w = quantized_4bit(&ffi::astype(&seeded(&[1024, k], 7), crate::dtype::BFLOAT16)); + assert!( + !prefill_dense_gemm_eligible(&batched_x, &w, 1), + "a batch axis above 1 must stay on qmm on ROCm" + ); + } + /// Pin [`max_abs_attention_score`] against products computed by hand, on /// both sides of [`F16_MAX`]. /// diff --git a/src/lib/mlxcel-core/src/lib.rs b/src/lib/mlxcel-core/src/lib.rs index 40ff608ef..daa18b0cd 100644 --- a/src/lib/mlxcel-core/src/lib.rs +++ b/src/lib/mlxcel-core/src/lib.rs @@ -2344,6 +2344,22 @@ mod ffi { /// its first port. fn bitlinear_kernel_available() -> bool; + /// Whether `quantized_matmul` for this transposed affine projection + /// runs the same dense GEMM as `dequantize` + `matmul`, so the two + /// return the same bytes (lablup/mlxcel#2081). Reads the + /// QuantizedMatmul route on ROCm; true on every other backend, whose + /// rule is the output-tile count in + /// [`crate::layers::prefill_dense_gemm_eligible`]. `biases` may be + /// null. + unsafe fn quantized_matmul_matches_dense_gemm( + x: &MlxArray, + weight: &MlxArray, + scales: &MlxArray, + biases: *const MlxArray, + group_size: i32, + bits: i32, + ) -> bool; + /// True when this backend has a fused affine 1-bit kernel port /// (Metal). Other backends run 1-bit weights through the dequantize /// graph (issue #1370).