From f686c4c21450904c9e3c1678f930da6230992e7d Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Sat, 26 Sep 2026 16:42:13 +0200 Subject: [PATCH 1/3] gather-qmm --- mlx/backend/metal/kernels/fp_quantized.h | 164 +++++------- mlx/backend/metal/kernels/fp_quantized_nax.h | 251 +++++++++---------- mlx/backend/metal/kernels/quantized.h | 162 +++++------- mlx/backend/metal/kernels/quantized_nax.h | 222 ++++++++-------- mlx/backend/metal/quantized.cpp | 26 +- 5 files changed, 347 insertions(+), 478 deletions(-) diff --git a/mlx/backend/metal/kernels/fp_quantized.h b/mlx/backend/metal/kernels/fp_quantized.h index 52f933e55f..b137a6d473 100644 --- a/mlx/backend/metal/kernels/fp_quantized.h +++ b/mlx/backend/metal/kernels/fp_quantized.h @@ -2050,11 +2050,12 @@ template < const device uint32_t* w, const device uint8_t* scales, const device float* global_scale, - const device uint32_t* indices, + const device int32_t* offsets, device T* y, const constant int& M, const constant int& N, const constant int& K, + const constant int& num_groups, uint3 tid [[threadgroup_position_in_grid]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint simd_lane_id [[thread_index_in_simdgroup]]) { @@ -2099,13 +2100,18 @@ template < const int K_it = K / BK; const size_t stride_w = transpose ? N * K_w : K * N_w; const size_t stride_s = transpose ? N * K_g : K * N_g; - const int y_row = tid.y * BM; + int y_row; + int group; + short tgp_bm; + if (!schedule_row_tile( + offsets, num_groups, M, tid.y, simd_lane_id, y_row, group, tgp_bm)) { + return; + } const int y_col = tid.x * BN; const size_t y_row_long = size_t(y_row); const size_t y_col_long = size_t(y_col); // Prepare threadgroup bounds - const short tgp_bm = align_M ? BM : short(min(BM, M - y_row)); const short tgp_bn = align_N ? BN : short(min(BN, N - y_col)); // Calculate the final tiles in the case that K is not aligned @@ -2121,113 +2127,61 @@ template < wl += transpose ? y_col_long * K_w : y_col * bytes_per_pack / pack_factor; scales += transpose ? y_col_long * K_g : y_col / group_size; - // Do as many matmuls as necessary - uint32_t index; - short offset; - uint32_t index_next = indices[y_row]; - short offset_next = 0; - int n = 0; - while (n < tgp_bm) { - n++; - offset = offset_next; - index = index_next; - offset_next = tgp_bm; - for (; n < tgp_bm; n++) { - if (indices[y_row + n] != index) { - offset_next = n; - index_next = indices[y_row + n]; - break; - } - } - threadgroup_barrier(mem_flags::mem_none); - - // Prepare threadgroup mma operation - thread mma_t mma_op(simd_group_id, simd_lane_id); - - // Prepare threadgroup loading operations - thread loader_x_t loader_x(x, K, Xs, simd_group_id, simd_lane_id); - thread loader_w_t loader_w( - wl + index * stride_w, - scales + index * stride_s, - transpose ? K : N, - Ws, - simd_group_id, - simd_lane_id, - global_scale + index); - - // Matrices are all aligned check nothing - if (align_M && align_N) { - gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } + // Prepare threadgroup mma operation + thread mma_t mma_op(simd_group_id, simd_lane_id); - // Store results to device memory - if (offset_next - offset == BM) { - mma_op.store_result(y, N); - } else { - mma_op.store_result_slice( - y, N, short2(0, offset), short2(BN, offset_next)); - } - } else { - // Tile aligned so check outside of the hot loop - if ((align_M || tgp_bm == BM) && (align_N || tgp_bn == BN)) { - gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - - // Store results to device memory - if (offset_next - offset == BM) { - mma_op.store_result(y, N); - } else { - mma_op.store_result_slice( - y, N, short2(0, offset), short2(BN, offset_next)); - } - } + // Prepare threadgroup loading operations + thread loader_x_t loader_x(x, K, Xs, simd_group_id, simd_lane_id); + thread loader_w_t loader_w( + wl + group * stride_w, + scales + group * stride_s, + transpose ? K : N, + Ws, + simd_group_id, + simd_lane_id, + global_scale + group); + + // Tile aligned so check outside of the hot loop + if (tgp_bm == BM && (align_N || tgp_bn == BN)) { + gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); + } + mma_op.store_result(y, N); + } - // Tile partially aligned check rows - else if (align_N || tgp_bn == BN) { - gemm_loop_unaligned( - Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - mma_op.store_result_slice( - y, N, short2(0, offset), short2(BN, offset_next)); - } + // Tile partially aligned check rows + else if (align_N || tgp_bn == BN) { + gemm_loop_unaligned( + Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); + } + mma_op.store_result_safe(y, N, short2(BN, tgp_bm)); + } - // Tile partially aligned check cols - else if (align_M || tgp_bm == BM) { - gemm_loop_unaligned( - Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - mma_op.store_result_slice( - y, N, short2(0, offset), short2(tgp_bn, offset_next)); - } + // Tile partially aligned check cols + else if (tgp_bm == BM) { + gemm_loop_unaligned( + Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); + } + mma_op.store_result_safe(y, N, short2(tgp_bn, BM)); + } - // Nothing aligned so check both rows and cols - else { - gemm_loop_unaligned( - Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - mma_op.store_result_slice( - y, N, short2(0, offset), short2(tgp_bn, offset_next)); - } + // Nothing aligned so check both rows and cols + else { + gemm_loop_unaligned( + Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); } + mma_op.store_result_safe(y, N, short2(tgp_bn, tgp_bm)); } } diff --git a/mlx/backend/metal/kernels/fp_quantized_nax.h b/mlx/backend/metal/kernels/fp_quantized_nax.h index b934d91c6f..f843c99960 100644 --- a/mlx/backend/metal/kernels/fp_quantized_nax.h +++ b/mlx/backend/metal/kernels/fp_quantized_nax.h @@ -846,11 +846,12 @@ template < const device uint32_t* w, const device uint8_t* scales, const device float* global_scale, - const device uint32_t* indices, + const device int32_t* offsets, device T* y, const constant int& M, const constant int& N, const constant int& K, + const constant int& num_groups, uint3 tid [[threadgroup_position_in_grid]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint simd_lane_id [[thread_index_in_simdgroup]]) { @@ -880,13 +881,18 @@ template < const int K_it = K / BK; const size_t stride_w = transpose ? N * K_w : K * N_w; const size_t stride_s = transpose ? N * K_g : K * N_g; - const int y_row = tid.y * BM; + int y_row; + int group; + short tgp_bm; + if (!schedule_row_tile( + offsets, num_groups, M, tid.y, simd_lane_id, y_row, group, tgp_bm)) { + return; + } const int y_col = tid.x * BN; const size_t y_row_long = size_t(y_row); const size_t y_col_long = size_t(y_col); // Prepare threadgroup bounds - const short tgp_bm = align_M ? BM : short(min(BM, M - y_row)); const short tgp_bn = align_N ? BN : short(min(BN, N - y_col)); // Calculate the final tiles in the case that K is not aligned @@ -912,10 +918,10 @@ 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(int(SM), max(0, (M - (y_row + tm)))); + const short sgp_sm = short(clamp(int(tgp_bm) - tm, 0, int(SM))); const short sgp_sn = align_N ? SN : min(int(SN), max(0, (N - (y_col + tn)))); - const bool is_unaligned_sm = align_M ? false : (sgp_sm != SM); + const bool rows_in_bounds = y_row + tm + SM <= M; const bool is_unaligned_bn = align_N ? false : (tgp_bn != BN); constexpr short BR = transpose ? TN : TK; @@ -923,150 +929,123 @@ template < using AccumType = float; - // Do as many matmuls as necessary - uint32_t index; - short offset; - uint32_t index_next = indices[y_row]; - short offset_next = 0; - int n = 0; - while (n < tgp_bm) { - n++; - offset = offset_next; - index = index_next; - offset_next = tgp_bm; - for (; n < tgp_bm; n++) { - if (indices[y_row + n] != index) { - offset_next = n; - index_next = indices[y_row + n]; - break; - } - } - 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(); - - const device T* xn = x + tm * K; - - // Prepare threadgroup loading operations - thread loader_w_t loader_w( - wl + index * stride_w, - scales + index * stride_s, - transpose ? K : N, - Ws, - simd_group_id, - simd_lane_id, - global_scale + index); - - dispatch_bool(align_M || !is_unaligned_sm, [&](auto kAlignedM) { - dispatch_bool(align_N || !is_unaligned_bn, [&](auto kAlignedN) { - for (int k = 0; k < K_it; k++) { - threadgroup_barrier(mem_flags::mem_threadgroup); - if constexpr (kAlignedN.value) { - loader_w.load_unsafe(); - } else { - loader_w.load_safe( - transpose ? short2(BK, tgp_bn) : short2(tgp_bn, BK)); - } + const bool sg_active = sgp_sm > 0; - threadgroup_barrier(mem_flags::mem_threadgroup); - - 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 Dtile; + Dtile.clear(); - xn += BK; - loader_w.next(); + const device T* xn = x + tm * K; + + // Prepare threadgroup loading operations + thread loader_w_t loader_w( + wl + group * stride_w, + scales + group * stride_s, + transpose ? K : N, + Ws, + simd_group_id, + simd_lane_id, + global_scale + group); + + dispatch_bool(rows_in_bounds, [&](auto kAlignedM) { + dispatch_bool(align_N || !is_unaligned_bn, [&](auto kAlignedN) { + for (int k = 0; k < K_it; k++) { + threadgroup_barrier(mem_flags::mem_threadgroup); + if constexpr (kAlignedN.value) { + loader_w.load_unsafe(); + } else { + loader_w.load_safe( + transpose ? short2(BK, tgp_bn) : short2(tgp_bn, BK)); } - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - loader_w.load_safe(tile_w); - threadgroup_barrier(mem_flags::mem_threadgroup); - - 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, (k_remain - 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); + + 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; } } + xn += BK; + loader_w.next(); + } + + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + loader_w.load_safe(tile_w); threadgroup_barrier(mem_flags::mem_threadgroup); - // Store results to device memory - if constexpr (kAlignedN.value) { - if (m_lo_lim == 0 && m_hi_lim == SM) { - Dtile.store(y + tm * N + tn, N); - } else { - Dtile.store_slice( - y + tm * N + tn, N, short2(0, m_lo_lim), short2(SN, m_hi_lim)); + 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, (k_remain - 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); + + // Store results to device memory + if constexpr (kAlignedN.value) { + if (sgp_sm == SM) { + Dtile.store(y + tm * N + tn, N); } else { Dtile.store_slice( - y + tm * N + tn, - N, - short2(0, m_lo_lim), - short2(sgp_sn, m_hi_lim)); + y + tm * N + tn, N, short2(0, 0), short2(SN, sgp_sm)); } - }); + } else { + Dtile.store_slice( + y + tm * N + tn, N, short2(0, 0), short2(sgp_sn, sgp_sm)); + } }); - } + }); } diff --git a/mlx/backend/metal/kernels/quantized.h b/mlx/backend/metal/kernels/quantized.h index 9d9ce368c9..be1ebff597 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -2438,11 +2438,12 @@ template < const device uint32_t* w [[buffer(1)]], const device T* scales [[buffer(2)]], const device T* biases [[buffer(3)]], - const device uint32_t* indices [[buffer(4)]], + const device int32_t* offsets [[buffer(4)]], device T* y [[buffer(5)]], const constant int& M [[buffer(6)]], const constant int& N [[buffer(7)]], const constant int& K [[buffer(8)]], + const constant int& num_groups [[buffer(9)]], uint3 tid [[threadgroup_position_in_grid]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint simd_lane_id [[thread_index_in_simdgroup]]) { @@ -2486,13 +2487,18 @@ template < const int K_it = K / BK; const size_t stride_w = transpose ? N * K_w : K * N_w; const size_t stride_s = transpose ? N * K_g : K * N_g; - const int y_row = tid.y * BM; + int y_row; + int group; + short tgp_bm; + if (!schedule_row_tile( + offsets, num_groups, M, tid.y, simd_lane_id, y_row, group, tgp_bm)) { + return; + } const int y_col = tid.x * BN; const size_t y_row_long = size_t(y_row); const size_t y_col_long = size_t(y_col); // Prepare threadgroup bounds - const short tgp_bm = align_M ? BM : short(min(BM, M - y_row)); const short tgp_bn = align_N ? BN : short(min(BN, N - y_col)); // Calculate the final tiles in the case that K is not aligned @@ -2509,113 +2515,61 @@ template < scales += transpose ? y_col_long * K_g : y_col / group_size; biases += transpose ? y_col_long * K_g : y_col / group_size; - // Do as many matmuls as necessary - uint32_t index; - short offset; - uint32_t index_next = indices[y_row]; - short offset_next = 0; - int n = 0; - while (n < tgp_bm) { - n++; - offset = offset_next; - index = index_next; - offset_next = tgp_bm; - for (; n < tgp_bm; n++) { - if (indices[y_row + n] != index) { - offset_next = n; - index_next = indices[y_row + n]; - break; - } - } - threadgroup_barrier(mem_flags::mem_none); - - // Prepare threadgroup mma operation - thread mma_t mma_op(simd_group_id, simd_lane_id); - - // Prepare threadgroup loading operations - thread loader_x_t loader_x(x, K, Xs, simd_group_id, simd_lane_id); - thread loader_w_t loader_w( - wl + index * stride_w, - scales + index * stride_s, - biases + index * stride_s, - transpose ? K : N, - Ws, - simd_group_id, - simd_lane_id); - - // Matrices are all aligned check nothing - if (align_M && align_N) { - gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } + // Prepare threadgroup mma operation + thread mma_t mma_op(simd_group_id, simd_lane_id); - // Store results to device memory - if (offset_next - offset == BM) { - mma_op.store_result(y, N); - } else { - mma_op.store_result_slice( - y, N, short2(0, offset), short2(BN, offset_next)); - } - } else { - // Tile aligned so check outside of the hot loop - if ((align_M || tgp_bm == BM) && (align_N || tgp_bn == BN)) { - gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } + // Prepare threadgroup loading operations + thread loader_x_t loader_x(x, K, Xs, simd_group_id, simd_lane_id); + thread loader_w_t loader_w( + wl + group * stride_w, + scales + group * stride_s, + biases + group * stride_s, + transpose ? K : N, + Ws, + simd_group_id, + simd_lane_id); - // Store results to device memory - if (offset_next - offset == BM) { - mma_op.store_result(y, N); - } else { - mma_op.store_result_slice( - y, N, short2(0, offset), short2(BN, offset_next)); - } - } + // Tile aligned so check outside of the hot loop + if (tgp_bm == BM && (align_N || tgp_bn == BN)) { + gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); + } + mma_op.store_result(y, N); + } - // Tile partially aligned check rows - else if (align_N || tgp_bn == BN) { - gemm_loop_unaligned( - Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - mma_op.store_result_slice( - y, N, short2(0, offset), short2(BN, offset_next)); - } + // Tile partially aligned check rows + else if (align_N || tgp_bn == BN) { + gemm_loop_unaligned( + Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); + } + mma_op.store_result_safe(y, N, short2(BN, tgp_bm)); + } - // Tile partially aligned check cols - else if (align_M || tgp_bm == BM) { - gemm_loop_unaligned( - Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - mma_op.store_result_slice( - y, N, short2(0, offset), short2(tgp_bn, offset_next)); - } + // Tile partially aligned check cols + else if (tgp_bm == BM) { + gemm_loop_unaligned( + Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); + } + mma_op.store_result_safe(y, N, short2(tgp_bn, BM)); + } - // Nothing aligned so check both rows and cols - else { - gemm_loop_unaligned( - Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - gemm_loop_finalize( - Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); - } - mma_op.store_result_slice( - y, N, short2(0, offset), short2(tgp_bn, offset_next)); - } + // Nothing aligned so check both rows and cols + else { + gemm_loop_unaligned( + Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK); + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w); } + mma_op.store_result_safe(y, N, short2(tgp_bn, tgp_bm)); } } diff --git a/mlx/backend/metal/kernels/quantized_nax.h b/mlx/backend/metal/kernels/quantized_nax.h index 54a37482b6..a347749927 100644 --- a/mlx/backend/metal/kernels/quantized_nax.h +++ b/mlx/backend/metal/kernels/quantized_nax.h @@ -1465,11 +1465,12 @@ template < const device uint32_t* w [[buffer(1)]], const device T* scales [[buffer(2)]], const device T* biases [[buffer(3)]], - const device uint32_t* indices [[buffer(4)]], + const device int32_t* offsets [[buffer(4)]], device T* y [[buffer(5)]], const constant int& M [[buffer(6)]], const constant int& N [[buffer(7)]], const constant int& K [[buffer(8)]], + const constant int& num_groups [[buffer(9)]], uint3 tid [[threadgroup_position_in_grid]], uint simd_group_id [[simdgroup_index_in_threadgroup]], uint simd_lane_id [[thread_index_in_simdgroup]]) { @@ -1498,13 +1499,18 @@ template < const int K_it = K / BK; const size_t stride_w = transpose ? N * K_w : K * N_w; const size_t stride_s = transpose ? N * K_g : K * N_g; - const int y_row = tid.y * BM; + int y_row; + int group; + short tgp_bm; + if (!schedule_row_tile( + offsets, num_groups, M, tid.y, simd_lane_id, y_row, group, tgp_bm)) { + return; + } const int y_col = tid.x * BN; const size_t y_row_long = size_t(y_row); const size_t y_col_long = size_t(y_col); // Prepare threadgroup bounds - const short tgp_bm = align_M ? BM : short(min(BM, M - y_row)); const short tgp_bn = align_N ? BN : short(min(BN, N - y_col)); // Calculate the final tiles in the case that K is not aligned @@ -1531,11 +1537,11 @@ 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(int(SM), max(0, M - (y_row + tm))); + const short sgp_sm = short(clamp(int(tgp_bm) - tm, 0, int(SM))); const short sgp_sn = align_N ? SN : min(SN, short(max(0, (N - (y_col + tn))))); - const bool is_unaligned_sm = align_M ? false : (sgp_sm != SM); + const bool rows_in_bounds = y_row + tm + SM <= M; const bool is_unaligned_bn = align_N ? false : (tgp_bn != BN); constexpr short BR = transpose ? TN : TK; @@ -1543,145 +1549,119 @@ template < using AccumType = float; - // Do as many matmuls as necessary - uint32_t index; - short offset; - uint32_t index_next = indices[y_row]; - short offset_next = 0; - int n = 0; - while (n < tgp_bm) { - n++; - offset = offset_next; - index = index_next; - offset_next = tgp_bm; - for (; n < tgp_bm; n++) { - if (indices[y_row + n] != index) { - offset_next = n; - index_next = indices[y_row + n]; - break; - } - } - 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(); - - const device T* xn = x + tm * K; - - // Prepare threadgroup loading operations - thread loader_w_t loader_w( - wl + index * stride_w, - scales + index * stride_s, - biases + index * stride_s, - transpose ? K : N, - Ws, - simd_group_id, - simd_lane_id); - - dispatch_bool(align_M || !is_unaligned_sm, [&](auto kAlignedM) { - dispatch_bool(align_N || !is_unaligned_bn, [&](auto kAlignedN) { - for (int k = 0; k < K_it; k++) { - threadgroup_barrier(mem_flags::mem_threadgroup); - if constexpr (kAlignedN.value) { - loader_w.load_unsafe(); - } else { - loader_w.load_safe( - transpose ? short2(BK, tgp_bn) : short2(tgp_bn, BK)); - } + const bool sg_active = sgp_sm > 0; - threadgroup_barrier(mem_flags::mem_threadgroup); + NAXTile Dtile; + Dtile.clear(); - STEEL_PRAGMA_NO_UNROLL - for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { - NAXTile Atile; - NAXTile Btile; + const device T* xn = x + tm * K; + + // Prepare threadgroup loading operations + thread loader_w_t loader_w( + wl + group * stride_w, + scales + group * stride_s, + biases + group * stride_s, + transpose ? K : N, + Ws, + simd_group_id, + simd_lane_id); + + dispatch_bool(rows_in_bounds, [&](auto kAlignedM) { + dispatch_bool(align_N || !is_unaligned_bn, [&](auto kAlignedN) { + for (int k = 0; k < K_it; k++) { + threadgroup_barrier(mem_flags::mem_threadgroup); + if constexpr (kAlignedN.value) { + loader_w.load_unsafe(); + } else { + loader_w.load_safe( + transpose ? short2(BK, tgp_bn) : short2(tgp_bn, BK)); + } - volatile int compiler_barrier; + threadgroup_barrier(mem_flags::mem_threadgroup); - if constexpr (kAlignedM.value) { - Atile.load(xn + kk1, K); - } else { - Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); - } + STEEL_PRAGMA_NO_UNROLL + for (int kk1 = 0; kk1 < BK; kk1 += SK) { + if (sg_active) { + NAXTile Atile; + NAXTile Btile; - if constexpr (transpose) { - Btile.template load(Ws + tn * BK_padded + kk1); - } else { - Btile.template load(Ws + tn + kk1 * BN_padded); - } + volatile int compiler_barrier; - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant{}, - Btile, - metal::bool_constant{}); + if constexpr (kAlignedM.value) { + Atile.load(xn + kk1, K); + } else { + Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); + } - (void)compiler_barrier; + if constexpr (transpose) { + Btile.template load(Ws + tn * BK_padded + kk1); + } else { + Btile.template load(Ws + tn + kk1 * BN_padded); } - } - xn += BK; - loader_w.next(); - } + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant{}, + Btile, + metal::bool_constant{}); - if (!align_K) { - threadgroup_barrier(mem_flags::mem_threadgroup); - loader_w.load_safe(tile_w); - threadgroup_barrier(mem_flags::mem_threadgroup); + (void)compiler_barrier; + } + } - STEEL_PRAGMA_NO_UNROLL - for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { - NAXTile Atile; - NAXTile Btile; + xn += BK; + loader_w.next(); + } - volatile int compiler_barrier; + if (!align_K) { + threadgroup_barrier(mem_flags::mem_threadgroup); + loader_w.load_safe(tile_w); + threadgroup_barrier(mem_flags::mem_threadgroup); - const short psk = min(int(SK), max(0, (k_remain - kk1))); - Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); + STEEL_PRAGMA_NO_UNROLL + for (int kk1 = 0; kk1 < BK; kk1 += SK) { + if (sg_active) { + NAXTile Atile; + NAXTile Btile; - if constexpr (transpose) { - Btile.template load(Ws + tn * BK_padded + kk1); - } else { - Btile.template load(Ws + tn + kk1 * BN_padded); - } + volatile int compiler_barrier; - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant{}, - Btile, - metal::bool_constant{}); + const short psk = min(int(SK), max(0, (k_remain - kk1))); + Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); - (void)compiler_barrier; + 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); + threadgroup_barrier(mem_flags::mem_threadgroup); - // Store results to device memory - if constexpr (kAlignedN.value) { - if (m_lo_lim == 0 && m_hi_lim == SM) { - Dtile.store(y + tm * N + tn, N); - } else { - Dtile.store_slice( - y + tm * N + tn, N, short2(0, m_lo_lim), short2(SN, m_hi_lim)); - } + // Store results to device memory + if constexpr (kAlignedN.value) { + if (sgp_sm == SM) { + Dtile.store(y + tm * N + tn, N); } else { Dtile.store_slice( - y + tm * N + tn, - N, - short2(0, m_lo_lim), - short2(sgp_sn, m_hi_lim)); + y + tm * N + tn, N, short2(0, 0), short2(SN, sgp_sm)); } - }); + } else { + Dtile.store_slice( + y + tm * N + tn, N, short2(0, 0), short2(sgp_sn, sgp_sm)); + } }); - } + }); } diff --git a/mlx/backend/metal/quantized.cpp b/mlx/backend/metal/quantized.cpp index 55608183dc..bc953ff016 100644 --- a/mlx/backend/metal/quantized.cpp +++ b/mlx/backend/metal/quantized.cpp @@ -6,6 +6,7 @@ #include "mlx/backend/gpu/copy.h" #include "mlx/backend/metal/device.h" #include "mlx/backend/metal/kernels.h" +#include "mlx/backend/metal/matmul.h" #include "mlx/backend/metal/reduce.h" #include "mlx/backend/metal/unary.h" #include "mlx/backend/metal/utils.h" @@ -1565,7 +1566,6 @@ void gather_qmm_rhs_nax( int bn = 64, bk = 64; int wm = 2, wn = 2; - const bool align_M = (M % bm) == 0; const bool align_N = (N % bn) == 0; const bool align_K = (K % bk) == 0; @@ -1595,7 +1595,6 @@ void gather_qmm_rhs_nax( global_scale ? "_hgs" : ""); metal::MTLFCList func_consts = { - {&align_M, MTL::DataType::DataTypeBool, 200}, {&align_N, MTL::DataType::DataTypeBool, 201}, {&align_K, MTL::DataType::DataTypeBool, 202}, }; @@ -1606,13 +1605,13 @@ void gather_qmm_rhs_nax( concatenate( hash_name, kname, - "_align_M_", - align_M ? 't' : 'n', "_align_N_", align_N ? 't' : 'n', "_align_K_", align_K ? 't' : 'n'); + array offsets = gather_mm_offsets(indices, E, M, d, s); + // Get and set the kernel auto& compute_encoder = metal::get_command_encoder(s); auto kernel = get_gather_qmm_nax_kernel( @@ -1634,7 +1633,8 @@ void gather_qmm_rhs_nax( compute_encoder.set_compute_pipeline_state(kernel); MTL::Size group_dims(32, wn, wm); - MTL::Size grid_dims((N + bn - 1) / bn, (M + bm - 1) / bm, 1); + MTL::Size grid_dims( + (N + bn - 1) / bn, std::min(M, (M + bm - 1) / bm + E - 1), 1); compute_encoder.set_input_array(x, 0); compute_encoder.set_input_array(w, 1); @@ -1645,11 +1645,12 @@ void gather_qmm_rhs_nax( compute_encoder.set_input_array(*gs, 3); } int c = 4; - compute_encoder.set_input_array(indices, c++); + compute_encoder.set_input_array(offsets, c++); compute_encoder.set_output_array(out, c++); compute_encoder.set_bytes(M, c++); compute_encoder.set_bytes(N, c++); compute_encoder.set_bytes(K, c++); + compute_encoder.set_bytes(E, c++); compute_encoder.dispatch_threadgroups(grid_dims, group_dims); } @@ -1726,10 +1727,10 @@ void gather_qmm_rhs( } // TODO: Tune the block sizes + int E = w.size() / w.shape(-1) / w.shape(-2); int bm = 16, bn = 32, bk = 32; int wm = 1, wn = 2; - const bool align_M = (M % bm) == 0; const bool align_N = (N % bn) == 0; const bool align_K = (K % bk) == 0; @@ -1758,7 +1759,6 @@ void gather_qmm_rhs( global_scale ? "_hgs" : ""); metal::MTLFCList func_consts = { - {&align_M, MTL::DataType::DataTypeBool, 200}, {&align_N, MTL::DataType::DataTypeBool, 201}, {&align_K, MTL::DataType::DataTypeBool, 202}, }; @@ -1769,13 +1769,13 @@ void gather_qmm_rhs( concatenate( hash_name, kname, - "_align_M_", - align_M ? 't' : 'n', "_align_N_", align_N ? 't' : 'n', "_align_K_", align_K ? 't' : 'n'); + array offsets = gather_mm_offsets(indices, E, M, d, s); + // Get and set the kernel auto& compute_encoder = metal::get_command_encoder(s); auto kernel = get_gather_qmm_kernel( @@ -1797,7 +1797,8 @@ void gather_qmm_rhs( compute_encoder.set_compute_pipeline_state(kernel); MTL::Size group_dims(32, wn, wm); - MTL::Size grid_dims((N + bn - 1) / bn, (M + bm - 1) / bm, 1); + MTL::Size grid_dims( + (N + bn - 1) / bn, std::min(M, (M + bm - 1) / bm + E - 1), 1); compute_encoder.set_input_array(x, 0); compute_encoder.set_input_array(w, 1); @@ -1808,11 +1809,12 @@ void gather_qmm_rhs( compute_encoder.set_input_array(*gs, 3); } int c = 4; - compute_encoder.set_input_array(indices, c++); + compute_encoder.set_input_array(offsets, c++); compute_encoder.set_output_array(out, c++); compute_encoder.set_bytes(M, c++); compute_encoder.set_bytes(N, c++); compute_encoder.set_bytes(K, c++); + compute_encoder.set_bytes(E, c++); compute_encoder.dispatch_threadgroups(grid_dims, group_dims); } From f9256bf11f85837d740a9ef32819d3c99ba1e0b0 Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Sat, 26 Sep 2026 17:57:47 +0200 Subject: [PATCH 2/3] swiglu bench --- benchmarks/python/swiglu_qmm_bench.py | 72 +++++++++++++++++++++++++++ 1 file changed, 72 insertions(+) create mode 100644 benchmarks/python/swiglu_qmm_bench.py diff --git a/benchmarks/python/swiglu_qmm_bench.py b/benchmarks/python/swiglu_qmm_bench.py new file mode 100644 index 0000000000..2e80ec2fa7 --- /dev/null +++ b/benchmarks/python/swiglu_qmm_bench.py @@ -0,0 +1,72 @@ +# Copyright © 2026 Apple Inc. + +import mlx.core as mx +import mlx.nn as nn +from time_utils import time_fn + +SEQ_LENS = [512, 713, 1024, 2048, 4123, 8192, 8192 * 2] +MODES = ["affine", "nvfp4", "mxfp8"] + +# https://huggingface.co/zai-org/GLM-5.3-Flash-BF16/blob/main/config.json +# https://huggingface.co/zai-org/GLM-5.3/blob/main/config.json +# https://huggingface.co/Qwen/Qwen3.5-35B-A3B/blob/main/config.json +# https://huggingface.co/Qwen/Qwen3.5-122B-A10B/blob/main/config.json +# https://huggingface.co/Qwen/Qwen3.5-397B-A17B/blob/main/config.json +CONFIGS = { + "glm-5.3-flash": (4096, 2048, 288, 8), + "glm-5.3": (6144, 2048, 256, 8), + "qwen3.5-35b-a3b": (2048, 512, 256, 8), + "qwen3.5-122b-a10b": (3072, 1024, 256, 8), + "qwen3.5-397b-a17b": (4096, 1024, 512, 10), +} + + +def gather_sort(x, indices): + N, M = indices.shape + indices = indices.flatten() + order = mx.argsort(indices) + inv_order = mx.argsort(order) + return x.flatten(0, -3)[order // M], indices[order], inv_order + + +def scatter_unsort(x, inv_order, shape=None): + x = x[inv_order] + if shape is not None: + x = mx.unflatten(x, 0, shape) + return x + + +def time_gather_qmm(name, D, M, E, I, mode): + w1 = mx.random.normal((E, M, D), dtype=mx.bfloat16, scale=D**-0.5) + w2 = mx.random.normal((E, M, D), dtype=mx.bfloat16, scale=D**-0.5) + w3 = mx.random.normal((E, D, M), dtype=mx.bfloat16, scale=M**-0.5) + w1, w2, w3 = (mx.quantize(w, mode=mode) for w in (w1, w2, w3)) + mx.eval(w1, w2, w3) + + def gather_qmm(x, w1, w2, w3, indices, sort): + idx = indices + inv_order = None + if sort: + x, idx, inv_order = gather_sort(x, indices) + kwargs = dict(transpose=True, mode=mode, rhs_indices=idx, sorted_indices=sort) + gate = mx.gather_qmm(x, *w1, **kwargs) + up = mx.gather_qmm(x, *w2, **kwargs) + x = mx.gather_qmm(nn.silu(gate) * up, *w3, **kwargs) + if sort: + x = scatter_unsort(x, inv_order, indices.shape) + return x + + for N in SEQ_LENS: + x = mx.random.normal((N, 1, 1, D), dtype=mx.bfloat16) + scores = mx.random.uniform(shape=(N, E)) + indices = mx.argpartition(scores, E - I, axis=-1)[:, -I:].astype(mx.uint32) + mx.eval(x, indices) + + label = f"{name} {mode} N={N}" + time_fn(gather_qmm, x, w1, w2, w3, indices, True, msg=f"{label} swiglu") + + +if __name__ == "__main__": + for mode in MODES: + for name, config in CONFIGS.items(): + time_gather_qmm(name, *config, mode) From c835030e266bb5bca00e2e2cad753eb0951949b3 Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Sun, 27 Sep 2026 18:08:41 +0200 Subject: [PATCH 3/3] dont copy if row contiguous --- .gitignore | 1 + mlx/backend/metal/quantized.cpp | 3 +++ 2 files changed, 4 insertions(+) diff --git a/.gitignore b/.gitignore index 161d8d67eb..f5d859ff38 100644 --- a/.gitignore +++ b/.gitignore @@ -79,3 +79,4 @@ uv.lock .cache/ # vim *.swp +wip \ No newline at end of file diff --git a/mlx/backend/metal/quantized.cpp b/mlx/backend/metal/quantized.cpp index bc953ff016..97c418e872 100644 --- a/mlx/backend/metal/quantized.cpp +++ b/mlx/backend/metal/quantized.cpp @@ -67,6 +67,9 @@ inline array ensure_row_contiguous_matrix( const array& x, metal::Device& d, const Stream& s) { + if (x.flags().row_contiguous) { + return x; + } if (x.ndim() < 2) { if (x.strides()[0] == 1) { return x;