From adcbb91a9e238e0ede1b41bf032dbbfe15788728 Mon Sep 17 00:00:00 2001 From: Awni Hannun Date: Mon, 2 Feb 2026 18:54:01 -0800 Subject: [PATCH] Fix for NAX overflow. (#3092) --- .../kernels/steel/gemm/kernels/steel_gemm_fused_nax.h | 8 ++++++-- .../kernels/steel/gemm/kernels/steel_gemm_gather_nax.h | 8 ++++++-- .../kernels/steel/gemm/kernels/steel_gemm_splitk_nax.h | 8 ++++++-- 3 files changed, 18 insertions(+), 6 deletions(-) diff --git a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h index 8b7479ce..4ff92606 100644 --- a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h +++ b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h @@ -157,10 +157,14 @@ template < const short tm = SM * (simd_group_id / WN); const short tn = SN * (simd_group_id % WN); - const short sgp_sm = align_M ? SM : min(SM, short(params->M - (c_row + tm))); + const int sgp_sm_int = + align_M ? int(SM) : min(int(SM), params->M - (c_row + tm)); + const short sgp_sm = short(sgp_sm_int); const bool is_unaligned_sm = align_M ? false : (sgp_sm != SM); - const short sgp_sn = align_N ? SN : min(SN, short(params->N - (c_col + tn))); + const int sgp_sn_int = + align_N ? int(SN) : min(int(SN), params->N - (c_col + tn)); + const short sgp_sn = short(sgp_sn_int); const bool is_unaligned_sn = align_N ? false : (sgp_sn != SN); A += transpose_a ? tm : (tm * params->lda); diff --git a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_gather_nax.h b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_gather_nax.h index 0bd9a459..67cd7378 100644 --- a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_gather_nax.h +++ b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_gather_nax.h @@ -53,10 +53,14 @@ gather_mm_rhs_nax( const short tm = SM * (simd_group_id / WN); const short tn = SN * (simd_group_id % WN); - const short sgp_sm = align_M ? SM : min(SM, short(params->M - (c_row + tm))); + const int sgp_sm_int = + align_M ? int(SM) : min(int(SM), params->M - (c_row + tm)); + const short sgp_sm = short(sgp_sm_int); const bool is_unaligned_sm = align_M ? false : (sgp_sm != SM); - const short sgp_sn = align_N ? SN : min(SN, short(params->N - (c_col + tn))); + const int sgp_sn_int = + align_N ? int(SN) : min(int(SN), params->N - (c_col + tn)); + const short sgp_sn = short(sgp_sn_int); const bool is_unaligned_sn = align_N ? false : (sgp_sn != SN); A += transpose_a ? tm : (tm * params->lda); diff --git a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_splitk_nax.h b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_splitk_nax.h index ea748c19..1b6b8280 100644 --- a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_splitk_nax.h +++ b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_splitk_nax.h @@ -86,10 +86,14 @@ template < const short tm = SM * (simd_group_id / WN); const short tn = SN * (simd_group_id % WN); - const short sgp_sm = align_M ? SM : min(SM, short(params->M - (c_row + tm))); + const int sgp_sm_int = + align_M ? int(SM) : min(int(SM), params->M - (c_row + tm)); + const short sgp_sm = short(sgp_sm_int); const bool is_unaligned_sm = align_M ? false : (sgp_sm != SM); - const short sgp_sn = align_N ? SN : min(SN, short(params->N - (c_col + tn))); + const int sgp_sn_int = + align_N ? int(SN) : min(int(SN), params->N - (c_col + tn)); + const short sgp_sn = short(sgp_sn_int); const bool is_unaligned_sn = align_N ? false : (sgp_sn != SN); A += transpose_a ? tm : (tm * params->lda);