Skip to content

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

Description

@inureyes

Part of #1801

Problem / Background

Since #2085 (PR f9aefa39), bf16 affine GEMMs of 128 rows or more on RDNA 3.5 leave the fused WMMA kernel for dequantize plus hipBLASLt. That route inserts every dequantized weight into an LRU of 8 matrices or 256 MB. A forward pass touches far more than 8 distinct projections (Llama-class 8B: 32 layers x 7), so prefill cycles the LRU and never hits, yet the last entries, up to 256 MB of bf16 weights, stay alive afterwards through decode. LOCAL_FIXES.md item 29 and TECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.en.md:190 record this ("gets no hits across a model's projections in prefill ... the last entries (up to 256 MB) stay alive after prefill"); peak memory for it was not measured.

Current Behavior

All in src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/quantized/qmm.hip:

  • select_qmm_route (:889-934): for a WMMA-eligible shape, hipblaslt_is_faster = env != 1 && dequant && M >= wmma_qmm_max_m(d) && !fp8() returns QmmRoute::DequantGemm (the ceiling route); wmma_qmm_max_m (:878) is 128 on RDNA 3.5. Non-WMMA shapes reach the same enumerator through if (dequant) (:930).
  • QuantizedMatmul::eval_gpu (:3447): for DequantGemm it looks up the LRU keyed by weight, scales and bias pointers (DequantCacheKey, :936), and on a miss dequantizes and inserts (:3627-3680) when dequant_cache_capacity() (:981, default 8, MLX_ROCM_QMM_DEQUANT_CACHE_SIZE) and dequant_cache_max_bytes() (:997, default 256 MB, MLX_ROCM_QMM_DEQUANT_CACHE_MAX_BYTES) are non-zero. Each entry also holds references to the quantized sources.
  • The route predicate at :3442 (route == QmmRoute::DequantGemm && is_hipblaslt_available()) is the other reader of the enumerator.

Proposed Solution

  1. Add QmmRoute::DequantGemmAboveWmmaCeiling, returned instead of DequantGemm from the hipblaslt_is_faster branch only. eval_gpu treats it exactly like DequantGemm except that it neither looks up nor inserts into the LRU: it dequantizes into a fresh w_dequant that the command encoder frees after the GEMM. Update the predicate at :3442 to accept both enumerators.
  2. Leave the LRU unchanged for the other DequantGemm callers (f16 and non-WMMA shapes) in this issue, but count hits and misses there during the measurement below. If the f16 route also shows zero hits on every measured model, file a follow-up to turn its default off; do not change it here.
  3. Update LOCAL_FIXES.md item 29 to describe the new behavior and the measured numbers.

Rejected: shrinking MLX_ROCM_QMM_DEQUANT_CACHE_MAX_BYTES globally (keeps the useless insertions and still pins memory) and reordering the key (the miss is structural, not a hashing problem).

Implementation Notes

  • Reuse: the existing dequant and hipBLASLt code in the DequantGemm arm; only the cache block is skipped.
  • Edge cases: MLX_ROCM_WMMA_QMM=1 (env == 1) never takes the ceiling route, unchanged. MLX_ROCM_WMMA_QMM_MAX_M overrides still produce the new enumerator. The fp8 route is unaffected.
  • Constraints: logits must not change (same dequantized values, same GEMM); prefill throughput must not regress beyond run-to-run noise, since the cache never hit.

Acceptance Criteria

  • Ceiling-routed GEMMs never insert into the dequant LRU; after a bf16 prefill on gfx1151 the LRU is empty (assert via a debug counter or test hook exposed for the test).
  • A ROCm test runs one bf16 quantized matmul above the ceiling and checks the LRU entry count stays 0 while the output matches the existing reference; another below the ceiling path (f16, M >= MLX_ROCM_QMM_DEQUANT_M_THRESHOLD) still inserts.
  • Measured on gfx1151, before and after, with Gemma-3-4B-it-4bit (bf16, the model fix(rocm): route large bf16 qmm to hipBLASLt and gate dense prefill #2085 measured) at pp2048/tg128 and one f16 checkpoint: MLX peak memory, active_memory() after prefill, prefill and decode tok/s, medians of 3 through scripts/rocm_gpu_guard.sh, plus LRU hit and miss counts. Results added to docs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md.
  • Logit traces for Gemma 3 4B with scripts/compare_logit_traces.py --decided 2.0 before and after are byte-identical.

Verification

cargo test --release --features rocm --test rocm_qmm_dequant_cache -- --test-threads=1
make verify-rocm verify-rocm-overlay
scripts/rocm_gpu_guard.sh --idle-secs 60 -- scripts/bench_decode.sh models/mlx/gemma-3-4b-it-4bit

Metadata

Metadata

Assignees

No one assigned

    Labels

    area:coremlxcel-core: MLX FFI, primitives, KV cache, layersplatform:linuxLinux (CUDA / packaging) specificpriority:lowLow prioritystatus:doneCompletedtype:performancePerformance improvements

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions