Skip to content

update(rocm): skip the dequant LRU for GEMMs routed above the WMMA ceiling - #2236

Merged
inureyes merged 3 commits into
mainfrom
update/issue-2151-skip-dequant-lru
Oct 8, 2026
Merged

inureyes merged 3 commits into
mainfrom
update/issue-2151-skip-dequant-lru

Conversation

@inureyes

@inureyes inureyes commented Oct 8, 2026 •

Copy link
Copy Markdown
Member

Summary

bf16 affine GEMMs that select_qmm_route sends to dequantize + hipBLASLt because they are at or above the WMMA row ceiling (128 rows on RDNA 3.5) no longer go through the dequantized-weight LRU. A forward pass runs more distinct projections than the 8-entry, 256 MB cache holds, so it never hit, and its last entries stayed alive through decode.

  • New QmmRoute::DequantGemmAboveWmmaCeiling, returned only from the ceiling branch. QuantizedMatmul::eval_gpu runs the same dequantize and GEMM but into a temporary freed after the GEMM; quantized_matmul_runs_dequant_gemm accepts both enumerators. f16 and non-WMMA shapes keep the cache unchanged.
  • Cache counters (hits, misses, inserts, evictions, bypasses, entries, bytes) as relaxed atomics: rocm::dequant_cache_stats() in the overlay (stubbed in no_rocm.cpp), mlxcel_core::rocm_qmm_cache::dequant_cache_stats() through the bridge, and MLX_ROCM_QMM_DEQUANT_CACHE_STATS=1 to print them at exit.
  • tests/rocm_qmm_dequant_cache.rs, LOCAL_FIXES item 29, docs/environment-variables.md, and a dated section in docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md.

Measurements (gfx1151, pp2048/tg128, medians of 3 on 8753883, each run under rocm_gpu_guard.sh --idle-secs 60)

Model Arm Prefill tok/s Decode tok/s MLX peak Active after prefill Cache hits / misses / bypassed, entries at exit
gemma-3-4b-it-4bit (bf16) before 2872 (2819-2910) 62.75 (62.68-62.76) 4.52 GB 2.80 GB 0 / 340 / 0, 6 (230 MiB)
gemma-3-4b-it-4bit (bf16) after 2868 (2839-2945) 62.71 (62.55-62.73) 4.52 GB 2.56 GB 0 / 0 / 340, 0
Llama-3.1-8B-Instruct-4bit (f16) before 1136 (1106-1136) 36.63 (32.36-36.84) 6.92 GB 4.75 GB 0 / 320 / 0, 2 (224 MiB)
Llama-3.1-8B-Instruct-4bit (f16) after 1132 (1048-1143) 36.67 (36.58-36.92) 6.92 GB 4.75 GB 0 / 320 / 0, 2 (224 MiB)

Prefill and decode differences are inside the run-to-run range; no speedup is claimed. Before counts come from a counters-only build that routes like 8753883. One confirmation run per model on the rebased head (main ad84435) matched the after rows. The f16 cache also never hits, so #2232 tracks turning its default off.

Verification

All on gfx1151 (Radeon 8060S), ROCm 10.0.0 / HIP 7.15. The two test binaries and the overlay, dtype-key, port-dispatch and fmt gates were rerun on the final head (rebased onto main ae343d8):

  • cargo test --release --features rocm --test rocm_qmm_dequant_cache -- --test-threads=1: pass. With the ceiling branch returning DequantGemm again, the two bf16 ceiling cases fail (misses=1 inserts=1 bypasses=0 entries=1 on the first pass); the forced-WMMA and f16 cases still pass.
  • cargo test --release --features rocm --test rocm_qmm_env -- --test-threads=1: pass.
  • Logit traces (examples/logit_trace, Gemma 3 4B, 4 x 256-token chunks) before and after: byte-identical; scripts/compare_logit_traces.py --decided 2.0 reports 0 / 1024 disagreements.
  • make verify-rocm-overlay verify-versions verify-kernel-dtype-keys verify-kernel-port-dispatch verify-llama-compat verify-fmt, cargo test --test dead_doc_pointers, cargo clippy -p mlxcel-core --features rocm --lib --tests -- -D warnings, cargo clippy -p mlxcel --features rocm --test rocm_qmm_dequant_cache -- -D warnings (before the rebase): pass.
  • Full make verify-rocm with MLXCEL_ROCM_SMOKE_MODEL=models/mlx/Qwen3-0.6B-4bit at 79a8c9bb on 6dbfe7d7 (the later commit only adds the technical report): [verify-rocm] OK, 162 cargo test suites, 12191 passed, 0 failed, 399 ignored; the ROCm smoke generated 32 tokens on the GPU. An earlier run on the ae343d84 base was also green (12190 passed, 0 failed) before main gained update(rocm): route top-k with top-p up to vocab 152064 on ROCm #2218, which touches the bridge files.
  • Review: the orchestrator read the route, cache and bridge diff (relaxed atomic counters bumped once per GEMM on the dequant routes only, the exit report registered after the cache is constructed, the bridge returns zeros without MLXCEL_BRIDGE_ROCM_BACKEND); no CRITICAL or HIGH findings.

Not verified: Metal and CUDA (not available on this host). The change is in the ROCm overlay and the bridge function is stubbed to zeros off ROCm (#ifdef MLXCEL_BRIDGE_ROCM_BACKEND); the unit test stats_are_zero_off_rocm covers that stub but has not been built or run on those backends.

Closes #2151

@inureyes inureyes added status:review Under review type:performance Performance improvements priority:low Low priority area:core mlxcel-core: MLX FFI, primitives, KV cache, layers platform:linux Linux (CUDA / packaging) specific labels Oct 8, 2026
…iling

bf16 affine GEMMs at or above the fused WMMA kernel's row ceiling on RDNA 3.5 (128 rows by default) go to dequantize + hipBLASLt, and every dequantized weight went into the 8-entry, 256 MB LRU. A forward pass runs far more distinct projections than that, so prefill cycled the cache without a hit and its last entries, up to 256 MB of bf16 weights plus references to the quantized sources, stayed alive through decode.

select_qmm_route now returns QmmRoute::DequantGemmAboveWmmaCeiling from the ceiling branch only. eval_gpu runs it exactly like DequantGemm but dequantizes into a temporary that the encoder frees after the GEMM, and quantized_matmul_runs_dequant_gemm accepts both enumerators. f16 and non-WMMA shapes keep the cache unchanged.

The cache moved into one struct with relaxed atomic hit, miss, insert, eviction and bypass counters. rocm::dequant_cache_stats() (stubbed in no_rocm.cpp) reads them with the entry count and bytes, the bridge exposes them as mlxcel_core::rocm_qmm_cache::dequant_cache_stats(), and MLX_ROCM_QMM_DEQUANT_CACHE_STATS=1 prints them on stderr at exit.

tests/rocm_qmm_dequant_cache.rs runs one case per child process on gfx1151: bf16 above the default and an env ceiling bypasses the cache and matches dequantize + matmul bytewise, forced WMMA touches nothing, and f16 still misses, inserts and then hits. With the ceiling branch returning DequantGemm again both bf16 ceiling cases fail.

Refs #2151
Adds a dated section to docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md and updates LOCAL_FIXES item 29 with what the ceiling route change measured on gfx1151 at pp2048/tg128, medians of three guarded runs on 8753883 plus one confirmation run per model on the rebased head.

Gemma 3 4B 4-bit (bf16): the cache made 0 hits in 340 lookups and held 230 MiB after prefill; with the change it is bypassed 340 times and stays empty, active memory after prefill drops from 2.80 to 2.56 GB, MLX peak stays 4.52 GB, prefill and decode stay within run-to-run range, and the logit traces are byte-identical. Llama 3.1 8B 4-bit (f16) is unchanged and also gets 0 hits while holding 224 MiB; turning that default off is tracked in #2232.

Refs #2151
@inureyes
inureyes force-pushed the update/issue-2151-skip-dequant-lru branch from b41cbfb to 79a8c9b Compare October 8, 2026 11:38
Bilingual report for the dequant LRU skip above the WMMA ceiling, with the full make verify-rocm result: 162 suites, 12191 passed, 0 failed, 399 ignored.
@inureyes inureyes added status:done Completed and removed status:review Under review labels Oct 8, 2026
@inureyes
inureyes merged commit 2aa5211 into main Oct 8, 2026
26 of 27 checks passed
@inureyes
inureyes deleted the update/issue-2151-skip-dequant-lru branch October 8, 2026 12:11
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:core mlxcel-core: MLX FFI, primitives, KV cache, layers platform:linux Linux (CUDA / packaging) specific priority:low Low priority status:done Completed type:performance Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

perf(rocm): skip the dequant LRU for GEMMs routed above the WMMA ceiling

1 participant