You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
fix(rocm): range-check GEMM env integers and fix a stale HIP graph comment #2152
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.
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.
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.
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).
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.
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.
Part of #1801
Problem / Background
The ROCm overlay's integer env parsers in the GEMM paths accept any
longand thenstatic_cast<int>it, so a value that does not fit anintsilently wraps.MLX_ROCM_WMMA_QMM_MAX_M=4294967297becomes 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()andcapacity_from_env()), but these parsers predate it. Separately, a comment inuse_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, noerrno/INT_MAXcheck, no warning. Read forMLX_ROCM_QMM_DEQUANT_M_THRESHOLD(:792) andMLX_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 forMLX_ROCM_GEMM_{F32,BF16}{,_BATCHED}_SOLUTION_INDEX.dequant_cache_capacity()(qmm.hip:981) parsesMLX_ROCM_QMM_DEQUANT_CACHE_SIZEinline with the same cast.gpu_watchdog_seconds()(device.cpp:499-520) andcapacity_from_env()(lru_cache.h:26-46):errno = 0,strtolbase 10, whole string must be the number, rejectERANGE, 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
mlx/backend/rocm/env_int.hwithint 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 ownstaticcaching, so each warning prints once.parse_non_negative_int_envcopies andparse_positive_int_env; route every read listed above through the helper with these bounds:MLX_ROCM_QMM_DEQUANT_M_THRESHOLDandMLX_ROCM_WMMA_QMM_MAX_M1 toINT_MAX; the four solution indexes 0 toINT_MAX;MLX_ROCM_QMM_DEQUANT_CACHE_SIZE0 toINT_MAX(0 stays the documented off switch).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 byselect_qmm_routeinquantized/qmm.hip); drop the claim about which GEMM prefill uses.LOCAL_FIXES.mdentry 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
static_cast<int>(value)of an uncheckedstrtolresult remains inpatches-rocm/mlx/backend/rocm(grep -rn "parse_positive_int_env\|parse_non_negative_int_env" src/lib/mlx-cpp/patches-rocmreturns nothing).tests/rocm_qmm_env.rsruns one child process per value, astests/rocm_fft_cache_env.rsdoes, forMLX_ROCM_WMMA_QMM_MAX_Min {unset,128,4294967297,2147483648,-1,0,12abc, empty}, withMLX_ROCM_QMM_DEQUANT_M_THRESHOLD=1set 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_gemmfor 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 today4294967297wraps to a ceiling of 1 and makes it true.MLX_ROCM_QMM_DEQUANT_CACHE_SIZE=4294967304(warns, default 8) and=0(accepted, cache off, no warning).use_hip_graphs()comment no longer names a prefill GEMM.docs/environment-variables.mdentries for these variables state the accepted range and the warning, if they do not already.Scope added during implementation
atoi/strtolreads from the docs(rocm): document MLX_ROCM_GATHER_QMV_EXPERT_BATCHED and the other runtime MLX_ROCM_* variables #2180 filing routed through the same helper:MLX_ROCM_QMV_TILE_N(1 to 32, now read once),MLX_ROCM_GROUPED_PREFILL_MIN_B,MLX_ROCM_MOE_SEG_MINandMLX_GRIDX_MULT(1 toINT_MAX).Verification
cargo test --release --features rocm --test rocm_qmm_env -- --test-threads=1 make verify-rocm verify-rocm-overlay