Repository navigation
fix(rocm): run a decode batch's rows as one quantized product - #2253
Merged
Merged
Conversation
The serving harness measured per-request decode as (completion_tokens - 1) / (total_s - ttft), a window that also holds every tick the scheduler spends on other requests' prefill, so a slow batched step and interleaved prefill looked the same (#2156). With --metrics, bench_serving_concurrency.py now differences the batch scheduler counters around each level and prints decode steps, decode tokens, mixed steps, prefill chunks and tokens_predicted_seconds, plus mean occupancy (decode tokens over decode steps) and the mean decode step as one request sees it (tokens_predicted_seconds over decode tokens; the raw ratio over decode steps counts a step once per request and is printed beside it). The arithmetic lives in scripts/bench_serving_metrics.py, covered by tests/test_bench_serving_concurrency.py on canned /metrics text, including a check that the metric names still exist in src/server/routes/metrics.rs. benchmark_paged_decode_production.sh passes --metrics. Refs #2156
Batch-4 server decode of Meta-Llama-3.1-8B-Instruct-4bit on gfx1151 stepped at 166.7 ms against 31 ms single-stream (#2156). The batched forward feeds every projection a [B, 1, K] activation, and QuantizedMatmul::eval_gpu in the ROCm overlay read that as B one-row products: qmv_warp_shared_batched_kernel streamed the whole weight once per batch element and took 95% of the step. eval_gpu now folds the batch into the row count when the weight is unbatched and the activation row contiguous, as the Metal backend does, so the rows take qmv_wide_kernel in one pass. Folded bf16 rows would then have hit the fused WMMA kernel, which costs 7 to 8 times the one-row GEMV at 2 to 8 rows, so select_qmm_route keeps 4- and 8-bit GEMMs with biases of at most 8 rows off it unless MLX_ROCM_WMMA_QMM=1. Measured under rocm_gpu_guard.sh, medians of three: the server batch-4 ~1K step went from 166.7 to 65.9 ms (per-request decode 6.0 to 15.3 tok/s), single-stream decode 37.69 to 37.84 tok/s, Qwen3-0.6B bf16 batched steps 1.5 to 1.6x faster at batch 2 to 8, Gemma 3 4B unchanged. tests/rocm_qmm_batched_rows.rs checks the layouts return identical bytes and match an f32 reference; examples/qmm_batch_rows_probe.rs times them. LOCAL_FIXES item 43. Refs #2156
Results page for #2156 with the raw outputs, guard logs and session scripts: per-request decode tok/s, decode step ms, occupancy and the GPU-busy versus host-gap split at batch 1 and 4, ~1K and ~16K, before and after the batch-fold fix; the bf16 measurements behind the small-row route; and a verdict with its deciding measurement for each candidate in the issue. At ~16K the per-request rate is set by the other requests' chunked prefill, 78% of whose GPU time is flash SDPA; that is filed as #2251. The Metal reference ratio was not measured (no Metal host). Refs #2156
Bilingual report for the batched decode row fold on ROCm, with the full make verify-rocm result: 164 suites, 12195 passed, 0 failed, 403 ignored.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Batch-4 server decode on gfx1151 ran at 6 tok/s per request against 32 single-stream because the batched decode step itself was 5.3x the single-stream step. The batched forward feeds every projection a
[B, 1, K]activation, and the ROCm overlay'sQuantizedMatmul::eval_gpuread that asBone-row products:qmv_warp_shared_batched_kernelstreamed the whole weight once per batch element and took 95% of the step.qmm.hip: fold the batch into the row count when the weight is unbatched and the activation row contiguous (what the Metal backend does), so the rows takeqmv_wide_kernelin one pass. 4- and 8-bit GEMMs with biases of at most 8 rows stay off the fused WMMA kernel (7 to 8x slower than the GEMV there), so bf16 decode batches gain too. LOCAL_FIXES item 43.bench_serving_concurrency.py --metricsprints per level the batch counter deltas, mean occupancy and mean decode step ms (scripts/bench_serving_metrics.py,tests/test_bench_serving_concurrency.py).docs/benchmark_results/rocm-batched-decode-gfx1151-2026-10-08.mdwith raw data, the GPU-busy versus host-gap split, and a verdict per candidate.Measurements (Meta-Llama-3.1-8B-Instruct-4bit, f16 scales, medians of 3, guarded)
Single-stream
mlxcel-bench-decode(pp512/tg128): 37.69 before, 37.84 after. Model-level batched step, Qwen3-0.6B-4bit (bf16): 10.7/18.5/34.7 ms before, 6.9/12.0/22.2 after at batch 2/4/8; gemma-3-4b-it-4bit unchanged.Candidates: dequant GEMM route ruled out (no dequant or GEMM dispatch in the decode trace;
MLX_ROCM_QMM_DEQUANT_M_THRESHOLD=64no change; no LRU traffic at batch 8 either); batched qmv route confirmed and fixed; host scheduler work ruled out (3.1 to 6.7 ms host gaps per step); interleaved prefill ruled out at ~1K, confirmed at ~16K.Verification
On gfx1151 (Radeon 8060S, ROCm 7.15), every GPU run under
scripts/rocm_gpu_guard.sh, clean attempts only:cargo test --release --features rocm --test rocm_qmm_batched_rows -- --test-threads=1: pass (before and after the rebase onto 7a3fcc4)cargo test --release --features rocm --test rocm_qmm_dequant_cache -- --test-threads=1: passpython3 -m unittest discover -s tests -p 'test_*.py': 105 pass;pytest tests/test_bench_serving_concurrency.py: 12 passcargo clippy --features rocm --example qmm_batch_rows_probe --test rocm_qmm_batched_rows -- -D warnings: cleanmake verify-versions verify-kernel-dtype-keys verify-kernel-port-dispatch verify-llama-compat verify-fmt verify-rocm-overlay verify-binary-assets,cargo test --features rocm --test dead_doc_pointers: passFull
make verify-rocmwithMLXCEL_ROCM_SMOKE_MODEL=models/mlx/Qwen3-0.6B-4bitat474163e3on7a3fcc4a(the later commit only adds the technical report):[verify-rocm] OK, 164 cargo test suites, 12195 passed, 0 failed, 403 ignored; the ROCm smoke generated 32 tokens on the GPU.Review (orchestrator): the fold applies only to an unbatched weight with a row-contiguous activation whose batch count matches, and the output is allocated row contiguous, so
[batch, M, N]and[batch * M, N]share bytes;tests/rocm_qmm_batched_rows.rschecks the layouts byte for byte. Known gaps recorded in the report: the small-row route check does not consult the debug switchMLX_QMV_NO_WIDE, the test comment still names the fused WMMA route for its bf16 cases (they now take qmv), and batch sizes above 8 were not measured.Not verified
patches-rocm/.../qmm.hip), ROCm-gated tests and Python tooling.Closes #2156