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
- 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.
- 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.
- 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
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
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.mditem 29 andTECHNICAL_REPORTS/2085-rocm-bf16-qmm-route-hipblaslt-20260930.en.md:190record 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()returnsQmmRoute::DequantGemm(the ceiling route);wmma_qmm_max_m(:878) is 128 on RDNA 3.5. Non-WMMA shapes reach the same enumerator throughif (dequant)(:930).QuantizedMatmul::eval_gpu(:3447): forDequantGemmit looks up the LRU keyed by weight, scales and bias pointers (DequantCacheKey,:936), and on a miss dequantizes and inserts (:3627-3680) whendequant_cache_capacity()(:981, default 8,MLX_ROCM_QMM_DEQUANT_CACHE_SIZE) anddequant_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.:3442(route == QmmRoute::DequantGemm && is_hipblaslt_available()) is the other reader of the enumerator.Proposed Solution
QmmRoute::DequantGemmAboveWmmaCeiling, returned instead ofDequantGemmfrom thehipblaslt_is_fasterbranch only.eval_gputreats it exactly likeDequantGemmexcept that it neither looks up nor inserts into the LRU: it dequantizes into a freshw_dequantthat the command encoder frees after the GEMM. Update the predicate at:3442to accept both enumerators.DequantGemmcallers (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.LOCAL_FIXES.mditem 29 to describe the new behavior and the measured numbers.Rejected: shrinking
MLX_ROCM_QMM_DEQUANT_CACHE_MAX_BYTESglobally (keeps the useless insertions and still pins memory) and reordering the key (the miss is structural, not a hashing problem).Implementation Notes
DequantGemmarm; only the cache block is skipped.MLX_ROCM_WMMA_QMM=1(env == 1) never takes the ceiling route, unchanged.MLX_ROCM_WMMA_QMM_MAX_Moverrides still produce the new enumerator. The fp8 route is unaffected.Acceptance Criteria
M >= MLX_ROCM_QMM_DEQUANT_M_THRESHOLD) still inserts.MLX peak memory,active_memory()after prefill, prefill and decode tok/s, medians of 3 throughscripts/rocm_gpu_guard.sh, plus LRU hit and miss counts. Results added todocs/benchmark_results/rocm-bf16-qmm-route-gfx1151-2026-09-30.md.scripts/compare_logit_traces.py --decided 2.0before 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