From d362a93a53a874a79e1f77327ede400a04ab5481 Mon Sep 17 00:00:00 2001 From: XXXXRT666 <157766680+XXXXRT666@users.noreply.github.com> Date: Wed, 29 Jul 2026 00:10:06 +0800 Subject: [PATCH] Skip empty NAX GEMM output groups --- .../metal/kernels/steel/gemm/gemm_nax.h | 10 ++++++++++ .../steel/gemm/kernels/steel_gemm_fused_nax.h | 19 +++++++++++-------- 2 files changed, 21 insertions(+), 8 deletions(-) diff --git a/mlx/backend/metal/kernels/steel/gemm/gemm_nax.h b/mlx/backend/metal/kernels/steel/gemm/gemm_nax.h index 40066be7a4..c42aa1306f 100644 --- a/mlx/backend/metal/kernels/steel/gemm/gemm_nax.h +++ b/mlx/backend/metal/kernels/steel/gemm/gemm_nax.h @@ -45,11 +45,17 @@ auto gemm_loop( NAXTile Dtile; Dtile.clear(); + const bool has_output = sgp_sm > 0 && sgp_sn > 0; + int gemm_k_iterations_ = gemm_k_iterations_aligned; STEEL_PRAGMA_NO_UNROLL for (int kk0 = 0; kk0 < gemm_k_iterations_; kk0++) { threadgroup_barrier(mem_flags::mem_none); + if constexpr (!kAlignedM || !kAlignedN) { + if (!has_output) + continue; + } STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { @@ -94,6 +100,10 @@ auto gemm_loop( if constexpr (!kAlignedK) { simdgroup_barrier(mem_flags::mem_none); + if constexpr (!kAlignedM || !kAlignedN) { + if (!has_output) + return Dtile; + } const short rem_bk = K - gemm_k_iterations_ * BK; 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 80d90333d7..8267707ade 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 @@ -200,14 +200,17 @@ template < params->gemm_k_iterations_aligned, sgp_sm, sgp_sn); - if (use_out_source) { - gemm_epilogue( - Dtile, C, params, addmm_params, sgp_sm, sgp_sn); - } - if constexpr (kAlignedM && kAlignedN) { - Dtile.store(D, int(params->ldd)); - } else { - Dtile.store_safe(D, int(params->ldd), short2(sgp_sn, sgp_sm)); + if ((kAlignedM.value || sgp_sm > 0) && + (kAlignedN.value || sgp_sn > 0)) { + if (use_out_source) { + gemm_epilogue( + Dtile, C, params, addmm_params, sgp_sm, sgp_sn); + } + if constexpr (kAlignedM && kAlignedN) { + Dtile.store(D, int(params->ldd)); + } else { + Dtile.store_safe(D, int(params->ldd), short2(sgp_sn, sgp_sm)); + } } }); });