Skip to content

fix(rocm): run a decode batch's rows as one quantized product - #2253

Merged
inureyes merged 4 commits into
mainfrom
fix/issue-2156-batched-decode-gap
Oct 8, 2026
Merged

inureyes merged 4 commits into
mainfrom
fix/issue-2156-batched-decode-gap

Conversation

@inureyes

@inureyes inureyes commented Oct 8, 2026 •

Copy link
Copy Markdown
Member

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's QuantizedMatmul::eval_gpu 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.

  • 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 take qmv_wide_kernel in 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 --metrics prints per level the batch counter deltas, mean occupancy and mean decode step ms (scripts/bench_serving_metrics.py, tests/test_bench_serving_concurrency.py).
  • Results page docs/benchmark_results/rocm-batched-decode-gfx1151-2026-10-08.md with raw data, the GPU-busy versus host-gap split, and a verdict per candidate.
  • The ~16K figure is set by the other requests' chunked prefill, 78% flash SDPA: filed as perf(rocm): flash SDPA dominates long-context chunked prefill on gfx1151 #2251.

Measurements (Meta-Llama-3.1-8B-Instruct-4bit, f16 scales, medians of 3, guarded)

Case tok/s per request, before after decode step ms, before after
batch 1, ~1K 32.0 32.0 31.3 31.2
batch 4, ~1K 6.0 15.3 166.7 65.9
batch 1, ~16K 24.6 24.3 40.8 41.2
batch 4, ~16K 5.6 7.1 383.1 364.0

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=64 no 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: pass
  • python3 -m unittest discover -s tests -p 'test_*.py': 105 pass; pytest tests/test_bench_serving_concurrency.py: 12 pass
  • cargo clippy --features rocm --example qmm_batch_rows_probe --test rocm_qmm_batched_rows -- -D warnings: clean
  • make 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: pass
  • Server matrix, probes and traces: commands and scripts in the results page's data directory

Full make verify-rocm with MLXCEL_ROCM_SMOKE_MODEL=models/mlx/Qwen3-0.6B-4bit at 474163e3 on 7a3fcc4a (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.rs checks the layouts byte for byte. Known gaps recorded in the report: the small-row route check does not consult the debug switch MLX_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

  • Metal reference ratio (issue step 3): no Metal host; recorded as not measured on the page.
  • Metal and CUDA builds: not available on this host. The change is confined to the ROCm overlay (patches-rocm/.../qmm.hip), ROCm-gated tests and Python tooling.

Closes #2156

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
@inureyes inureyes added status:review Under review type:bug Bug fixes, error corrections, or issue resolutions priority:medium Medium priority area:inference Generation, sampling, decoding (incl. speculative, DRY) platform:linux Linux (CUDA / packaging) specific labels Oct 8, 2026
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.
@inureyes inureyes added status:done Completed and removed status:review Under review labels Oct 8, 2026
@inureyes
inureyes merged commit 02d67b6 into main Oct 8, 2026
25 of 27 checks passed
@inureyes
inureyes deleted the fix/issue-2156-batched-decode-gap branch October 8, 2026 23:16
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:inference Generation, sampling, decoding (incl. speculative, DRY) platform:linux Linux (CUDA / packaging) specific priority:medium Medium priority status:done Completed type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

perf(rocm): find why batched server decode is far below single-stream

1 participant