Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
203 changes: 203 additions & 0 deletions TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.en.md

Large diffs are not rendered by default.

203 changes: 203 additions & 0 deletions TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.ko.md

Large diffs are not rendered by default.

48 changes: 48 additions & 0 deletions docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md
Original file line number Diff line number Diff line change
@@ -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.
16 changes: 15 additions & 1 deletion docs/environment-variables.md
Original file line number Diff line number Diff line change
Expand Up @@ -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/<owner>/<name>` 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/<owner>/<name>` 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 <PATH>` on `mlxcel-server` / `mlxcel serve` since #1438 reserved `--models-dir` for b10621 router mode; still `--models-dir <PATH>` on the `download` / `list` / `rm` / `generate` subcommands), then `MLXCEL_MODELS_DIR`, then `${MLXCEL_CACHE_DIR:-$HOME/.cache/mlxcel}/models`. (`download --local-dir <PATH>` is separate: it writes the snapshot verbatim at that exact path.) |
Expand Down Expand Up @@ -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

Expand Down
1 change: 1 addition & 0 deletions docs/mlxcelverse/upstream/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading
Loading