Skip to content

fix(rocm): range-check GEMM env integers and fix a stale HIP graph comment #2152

Description

@inureyes

Part of #1801

Problem / Background

The ROCm overlay's integer env parsers in the GEMM paths accept any long and then static_cast<int> it, so a value that does not fit an int silently wraps. MLX_ROCM_WMMA_QMM_MAX_M=4294967297 becomes a ceiling of 1, which sends every bf16 GEMM of two or more rows to dequantize plus hipBLASLt with no message. #2073 settled the rule for these knobs (gpu_watchdog_seconds() and capacity_from_env()), but these parsers predate it. Separately, a comment in use_hip_graphs() has been stale since #2085.

Current Behavior

  • parse_positive_int_env (src/lib/mlx-cpp/patches-rocm/mlx/backend/rocm/quantized/qmm.hip:434, from the fork via feat(rocm): vendor the ROCm backend into mlxcelverse, add rocm feature #1818): strtol, rejects junk and <= 0, no errno/INT_MAX check, no warning. Read for MLX_ROCM_QMM_DEQUANT_M_THRESHOLD (:792) and MLX_ROCM_WMMA_QMM_MAX_M (:880, the RDNA 3.5 ceiling fix(rocm): route large bf16 qmm to hipBLASLt and gate dense prefill #2085 added).
  • parse_non_negative_int_env, the same shape with < 0, exists in three copies: qmm.hip:464, matmul.cpp:178, gemms/rocblas_gemm.cpp:38, read for MLX_ROCM_GEMM_{F32,BF16}{,_BATCHED}_SOLUTION_INDEX.
  • dequant_cache_capacity() (qmm.hip:981) parses MLX_ROCM_QMM_DEQUANT_CACHE_SIZE inline with the same cast.
  • The reference rule: gpu_watchdog_seconds() (device.cpp:499-520) and capacity_from_env() (lru_cache.h:26-46): errno = 0, strtol base 10, whole string must be the number, reject ERANGE, out-of-range and junk, print one stderr line [ROCm] ignoring invalid NAME="value" (expected ...); using the default N, use the default.
  • use_hip_graphs() (device.cpp:43-54) says "prefill uses the WMMA GEMM". Since fix(rocm): route large bf16 qmm to hipBLASLt and gate dense prefill #2085, bf16 prefill of 128 rows or more on RDNA 3.5 uses dequantize plus hipBLASLt, and other tiers or dtypes use other routes.

Proposed Solution

  1. Add one overlay header mlx/backend/rocm/env_int.h with int env_int_or_default(const char* name, int default_value, long min_value, long max_value, const char* expected) implementing the fix(rocm): validate MLX_ROCM_FFT_CACHE_SIZE (0 is UB, junk throws) #2073 rule exactly (unset or empty gives the default silently, anything else invalid warns once per call site and gives the default). The callers keep their own static caching, so each warning prints once.
  2. Delete the three parse_non_negative_int_env copies and parse_positive_int_env; route every read listed above through the helper with these bounds: MLX_ROCM_QMM_DEQUANT_M_THRESHOLD and MLX_ROCM_WMMA_QMM_MAX_M 1 to INT_MAX; the four solution indexes 0 to INT_MAX; MLX_ROCM_QMM_DEQUANT_CACHE_SIZE 0 to INT_MAX (0 stays the documented off switch).
  3. Rewrite the use_hip_graphs() comment: decode uses build-once capture/replay (decode_capture_*) and prefill does not go through this path either (its GEMM route is chosen by select_qmm_route in quantized/qmm.hip); drop the claim about which GEMM prefill uses.
  4. Add one LOCAL_FIXES.md entry covering the parsers and the comment.

parse_non_negative_size_t_env (qmm.hip:447, strtoull) is out of scope: it already cannot wrap into a smaller type.

Acceptance Criteria

  • No static_cast<int>(value) of an unchecked strtol result remains in patches-rocm/mlx/backend/rocm (grep -rn "parse_positive_int_env\|parse_non_negative_int_env" src/lib/mlx-cpp/patches-rocm returns nothing).
  • tests/rocm_qmm_env.rs runs one child process per value, as tests/rocm_fft_cache_env.rs does, for MLX_ROCM_WMMA_QMM_MAX_M in {unset, 128, 4294967297, 2147483648, -1, 0, 12abc, empty}, with MLX_ROCM_QMM_DEQUANT_M_THRESHOLD=1 set in every child (so the dequant route is eligible and only the ceiling decides), and asserts both the warning count on stderr and the route: ffi::quantized_matmul_matches_dense_gemm for a bf16 [64, 4096] x 4-bit [4096, 4096] g64 GEMM is false for every invalid value (the default ceiling, 128 on gfx1151, applies), where today 4294967297 wraps to a ceiling of 1 and makes it true.
  • The same test covers MLX_ROCM_QMM_DEQUANT_CACHE_SIZE=4294967304 (warns, default 8) and =0 (accepted, cache off, no warning).
  • The use_hip_graphs() comment no longer names a prefill GEMM.
  • docs/environment-variables.md entries for these variables state the accepted range and the warning, if they do not already.

Scope added during implementation

Verification

cargo test --release --features rocm --test rocm_qmm_env -- --test-threads=1
make verify-rocm verify-rocm-overlay

No activity

Activity on this issue will appear here.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    area:coremlxcel-core: MLX FFI, primitives, KV cache, layersplatform:linuxLinux (CUDA / packaging) specificpriority:lowLow prioritystatus:doneCompletedtype:bugBug fixes, error corrections, or issue resolutions

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions