From e9debbeaa0f36cdc722ee24e9ef36987c9590e33 Mon Sep 17 00:00:00 2001 From: jonathan308 Date: Tue, 25 Aug 2026 08:33:24 -0700 Subject: [PATCH] Revert "Skip unnecessary simdgroup computations for quantised MOE matmuls on NAX (#4352)" This reverts commit c7ff35d9714c78bcdf4620deb1d189f7ffb7c3b9. --- mlx/backend/metal/kernels/fp_quantized_nax.h | 109 +++++++++---------- mlx/backend/metal/kernels/quantized_nax.h | 93 ++++++++-------- 2 files changed, 96 insertions(+), 106 deletions(-) diff --git a/mlx/backend/metal/kernels/fp_quantized_nax.h b/mlx/backend/metal/kernels/fp_quantized_nax.h index 946bce7868..57712d9bf2 100644 --- a/mlx/backend/metal/kernels/fp_quantized_nax.h +++ b/mlx/backend/metal/kernels/fp_quantized_nax.h @@ -897,10 +897,6 @@ template < threadgroup_barrier(mem_flags::mem_none); // Prepare threadgroup mma operation - const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); - const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); - const bool sg_active = m_hi_lim > m_lo_lim; - NAXTile Dtile; Dtile.clear(); @@ -930,35 +926,33 @@ template < STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { - NAXTile Atile; - NAXTile Btile; - - volatile int compiler_barrier; - - if constexpr (kAlignedM.value) { - Atile.load(xn + kk1, K); - } else { - Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); - } - - if constexpr (transpose) { - Btile.template load( - Ws + tn * BK_padded + kk1); - } else { - Btile.template load( - Ws + tn + kk1 * BN_padded); - } - - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant{}, - Btile, - metal::bool_constant{}); - - (void)compiler_barrier; + NAXTile Atile; + NAXTile Btile; + + volatile int compiler_barrier; + + if constexpr (kAlignedM.value) { + Atile.load(xn + kk1, K); + } else { + Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); } + + if constexpr (transpose) { + Btile.template load( + Ws + tn * BK_padded + kk1); + } else { + Btile.template load( + Ws + tn + kk1 * BN_padded); + } + + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant{}, + Btile, + metal::bool_constant{}); + + (void)compiler_barrier; } xn += BK; @@ -972,37 +966,38 @@ template < STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { - NAXTile Atile; - NAXTile Btile; - - volatile int compiler_barrier; - - const short psk = min(int(SK), max(0, (BK - kk1))); - Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); - - if constexpr (transpose) { - Btile.template load( - Ws + tn * BK_padded + kk1); - } else { - Btile.template load( - Ws + tn + kk1 * BN_padded); - } - - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant{}, - Btile, - metal::bool_constant{}); - - (void)compiler_barrier; + NAXTile Atile; + NAXTile Btile; + + volatile int compiler_barrier; + + const short psk = min(int(SK), max(0, (BK - kk1))); + Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); + + if constexpr (transpose) { + Btile.template load( + Ws + tn * BK_padded + kk1); + } else { + Btile.template load( + Ws + tn + kk1 * BN_padded); } + + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant{}, + Btile, + metal::bool_constant{}); + + (void)compiler_barrier; } } threadgroup_barrier(mem_flags::mem_threadgroup); + const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); + const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); + // Store results to device memory if constexpr (kAlignedN.value) { if (m_lo_lim == 0 && m_hi_lim == SM) { diff --git a/mlx/backend/metal/kernels/quantized_nax.h b/mlx/backend/metal/kernels/quantized_nax.h index ed32eb59a7..db20c64390 100644 --- a/mlx/backend/metal/kernels/quantized_nax.h +++ b/mlx/backend/metal/kernels/quantized_nax.h @@ -1576,10 +1576,6 @@ template < } threadgroup_barrier(mem_flags::mem_none); - const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); - const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); - const bool sg_active = m_hi_lim > m_lo_lim; - NAXTile Dtile; Dtile.clear(); @@ -1610,33 +1606,31 @@ template < STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { - NAXTile Atile; - NAXTile Btile; - - volatile int compiler_barrier; - - if constexpr (kAlignedM.value) { - Atile.load(xn + kk1, K); - } else { - Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); - } - - if constexpr (transpose) { - Btile.template load(Ws + tn * BK_padded + kk1); - } else { - Btile.template load(Ws + tn + kk1 * BN_padded); - } - - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant{}, - Btile, - metal::bool_constant{}); - - (void)compiler_barrier; + NAXTile Atile; + NAXTile Btile; + + volatile int compiler_barrier; + + if constexpr (kAlignedM.value) { + Atile.load(xn + kk1, K); + } else { + Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); + } + + if constexpr (transpose) { + Btile.template load(Ws + tn * BK_padded + kk1); + } else { + Btile.template load(Ws + tn + kk1 * BN_padded); } + + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant{}, + Btile, + metal::bool_constant{}); + + (void)compiler_barrier; } xn += BK; @@ -1650,35 +1644,36 @@ template < STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { - NAXTile Atile; - NAXTile Btile; + NAXTile Atile; + NAXTile Btile; - volatile int compiler_barrier; + volatile int compiler_barrier; - const short psk = min(int(SK), max(0, (BK - kk1))); - Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); + const short psk = min(int(SK), max(0, (BK - kk1))); + Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); - if constexpr (transpose) { - Btile.template load(Ws + tn * BK_padded + kk1); - } else { - Btile.template load(Ws + tn + kk1 * BN_padded); - } + if constexpr (transpose) { + Btile.template load(Ws + tn * BK_padded + kk1); + } else { + Btile.template load(Ws + tn + kk1 * BN_padded); + } - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant{}, - Btile, - metal::bool_constant{}); + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant{}, + Btile, + metal::bool_constant{}); - (void)compiler_barrier; - } + (void)compiler_barrier; } } threadgroup_barrier(mem_flags::mem_threadgroup); + const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); + const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); + // Store results to device memory if constexpr (kAlignedN.value) { if (m_lo_lim == 0 && m_hi_lim == SM) {