From 2ff83a9e1235155ad7b81a11b07aed3c34966e6f Mon Sep 17 00:00:00 2001 From: Asher Feldman <59994+asher@users.noreply.github.com> Date: Sat, 1 Aug 2026 00:29:42 -0700 Subject: [PATCH] fix(metal): fix/silence all kernel build warnings --- .../mlx/backend/metal/kernels/kq_quantized.h | 133 ++++++------ .../backend/metal/kernels/kq_quantized_fp.h | 48 ++--- .../backend/metal/kernels/kq_quantized_iq.h | 192 +++++++++--------- .../metal/kernels/kq_quantized_kquants.h | 120 +++++------ .../metal/kernels/kq_quantized_legacy.h | 96 ++++----- .../backend/metal/kernels/kq_quantized_nax.h | 1 + metal/mlx/backend/metal/kernels/kq_sdpa.h | 1 - 7 files changed, 296 insertions(+), 295 deletions(-) diff --git a/metal/mlx/backend/metal/kernels/kq_quantized.h b/metal/mlx/backend/metal/kernels/kq_quantized.h index 1eb90c0..36f9841 100644 --- a/metal/mlx/backend/metal/kernels/kq_quantized.h +++ b/metal/mlx/backend/metal/kernels/kq_quantized.h @@ -1667,12 +1667,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1736,12 +1736,12 @@ template const device T* bias, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1832,12 +1832,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1895,12 +1895,12 @@ template const device T* bias, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2606,12 +2606,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2682,12 +2682,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3220,36 +3220,37 @@ METAL_FUNC void kq_gather_qmm_rhs_impl( // rows per expert fragment each row tile into per-expert segments that each // pay a full-tile mma K-loop, so utilization is ~1/segments-per-tile. The // dispatch picks the largest BM not much above the batch's rows-per-expert. -#define KQ_DEFINE_GATHER_QMM_RHS(CODEC, LOADER) \ - template \ - [[kernel]] void kq_##CODEC##_gather_qmm_rhs( \ - const device T* x [[buffer(0)]], \ - const device uint8_t* w [[buffer(1)]], \ - const device uint8_t* scales [[buffer(2)]], \ - const device uint32_t* indices [[buffer(3)]], \ - device T* y [[buffer(4)]], \ - const constant int& M [[buffer(5)]], \ - const constant int& N [[buffer(6)]], \ - const constant int& K [[buffer(7)]], \ - uint3 tid [[threadgroup_position_in_grid]], \ - uint simd_gid [[simdgroup_index_in_threadgroup]], \ - uint simd_lid [[thread_index_in_simdgroup]]) { \ - constexpr int BK = 32, BN = 64; \ - constexpr int BK_padded = (BK + 16 / sizeof(T)); \ - using LoaderW = LOADER< \ - T, \ - BN, \ - BK, \ - BK_padded, \ - /*reduction_dim=*/1, \ - /*tgp_size=*/2 * 2 * SIMD_SIZE>; \ - static_assert( \ - group_size == LoaderW::weights_per_block, \ - #CODEC " gather_qmm_rhs requires group_size == block size"); \ - threadgroup T Xs[BM * BK_padded]; \ - threadgroup T Ws[BN * BK_padded]; \ - kq_gather_qmm_rhs_impl( \ - x, w, indices, y, Xs, Ws, M, N, K, tid, simd_gid, simd_lid); \ +#define KQ_DEFINE_GATHER_QMM_RHS(CODEC, LOADER) \ + template \ + [[kernel]] void kq_##CODEC##_gather_qmm_rhs( \ + const device T* x [[buffer(0)]], \ + const device uint8_t* w [[buffer(1)]], \ + const device uint8_t* scales [[buffer(2)]], \ + const device uint32_t* indices [[buffer(3)]], \ + device T* y [[buffer(4)]], \ + const constant int& M [[buffer(5)]], \ + const constant int& N [[buffer(6)]], \ + const constant int& K [[buffer(7)]], \ + uint3 tid [[threadgroup_position_in_grid]], \ + uint simd_gid [[simdgroup_index_in_threadgroup]], \ + uint simd_lid [[thread_index_in_simdgroup]]) { \ + (void)scales; /* wire-format blocks are self-scaled; no separate scales */ \ + constexpr int BK = 32, BN = 64; \ + constexpr int BK_padded = (BK + 16 / sizeof(T)); \ + using LoaderW = LOADER< \ + T, \ + BN, \ + BK, \ + BK_padded, \ + /*reduction_dim=*/1, \ + /*tgp_size=*/2 * 2 * SIMD_SIZE>; \ + static_assert( \ + group_size == LoaderW::weights_per_block, \ + #CODEC " gather_qmm_rhs requires group_size == block size"); \ + threadgroup T Xs[BM * BK_padded]; \ + threadgroup T Ws[BN * BK_padded]; \ + kq_gather_qmm_rhs_impl( \ + x, w, indices, y, Xs, Ws, M, N, K, tid, simd_gid, simd_lid); \ } KQ_DEFINE_GATHER_QMM_RHS(q8_0, KqQ8_0BlockLoader) diff --git a/metal/mlx/backend/metal/kernels/kq_quantized_fp.h b/metal/mlx/backend/metal/kernels/kq_quantized_fp.h index 4cc3380..65b46a7 100644 --- a/metal/mlx/backend/metal/kernels/kq_quantized_fp.h +++ b/metal/mlx/backend/metal/kernels/kq_quantized_fp.h @@ -459,12 +459,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -519,12 +519,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -963,12 +963,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1023,12 +1023,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], diff --git a/metal/mlx/backend/metal/kernels/kq_quantized_iq.h b/metal/mlx/backend/metal/kernels/kq_quantized_iq.h index 40b872f..91781f6 100644 --- a/metal/mlx/backend/metal/kernels/kq_quantized_iq.h +++ b/metal/mlx/backend/metal/kernels/kq_quantized_iq.h @@ -1240,12 +1240,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1301,12 +1301,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1691,12 +1691,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1752,12 +1752,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2147,12 +2147,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2208,12 +2208,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2599,12 +2599,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2660,12 +2660,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3054,12 +3054,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3115,12 +3115,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3511,12 +3511,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3572,12 +3572,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3962,12 +3962,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -4023,12 +4023,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -4435,12 +4435,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -4496,12 +4496,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], diff --git a/metal/mlx/backend/metal/kernels/kq_quantized_kquants.h b/metal/mlx/backend/metal/kernels/kq_quantized_kquants.h index 214bf83..c5e499c 100644 --- a/metal/mlx/backend/metal/kernels/kq_quantized_kquants.h +++ b/metal/mlx/backend/metal/kernels/kq_quantized_kquants.h @@ -768,12 +768,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -844,12 +844,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1830,12 +1830,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1906,12 +1906,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2743,12 +2743,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2977,12 +2977,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3859,12 +3859,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -3935,12 +3935,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -4750,12 +4750,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -4826,12 +4826,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], diff --git a/metal/mlx/backend/metal/kernels/kq_quantized_legacy.h b/metal/mlx/backend/metal/kernels/kq_quantized_legacy.h index e0e63b9..f667ebe 100644 --- a/metal/mlx/backend/metal/kernels/kq_quantized_legacy.h +++ b/metal/mlx/backend/metal/kernels/kq_quantized_legacy.h @@ -598,12 +598,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -674,12 +674,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1359,12 +1359,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -1435,12 +1435,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2134,12 +2134,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2210,12 +2210,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2696,12 +2696,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], @@ -2756,12 +2756,12 @@ template device T* y, const constant int& in_vec_size, const constant int& out_vec_size, - const constant int& x_batch_ndims, - const constant int* x_shape, - const constant int64_t* x_strides, - const constant int& w_batch_ndims, - const constant int* w_shape, - const constant int64_t* w_strides, + const constant int& /* x_batch_ndims */, + const constant int* /* x_shape */, + const constant int64_t* /* x_strides */, + const constant int& /* w_batch_ndims */, + const constant int* /* w_shape */, + const constant int64_t* /* w_strides */, const constant int64_t* /* s_strides */, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], diff --git a/metal/mlx/backend/metal/kernels/kq_quantized_nax.h b/metal/mlx/backend/metal/kernels/kq_quantized_nax.h index 446bc90..0988bd3 100644 --- a/metal/mlx/backend/metal/kernels/kq_quantized_nax.h +++ b/metal/mlx/backend/metal/kernels/kq_quantized_nax.h @@ -3650,6 +3650,7 @@ METAL_FUNC void kq_gather_qmm_rhs_nax_tgp_impl( uint3 tid [[threadgroup_position_in_grid]], \ uint simd_group_id [[simdgroup_index_in_threadgroup]], \ uint simd_lane_id [[thread_index_in_simdgroup]]) { \ + (void)scales; /* wire-format blocks are self-scaled; no separate scales */ \ static_assert( \ group_size == GROUP_CONST, \ #codec " NAX kernel requires group_size=" #GROUP_CONST); \ diff --git a/metal/mlx/backend/metal/kernels/kq_sdpa.h b/metal/mlx/backend/metal/kernels/kq_sdpa.h index 9955e20..c31651a 100644 --- a/metal/mlx/backend/metal/kernels/kq_sdpa.h +++ b/metal/mlx/backend/metal/kernels/kq_sdpa.h @@ -245,7 +245,6 @@ template constexpr int D4 = D / 4; constexpr int NL = 32 / NE; // lanes per in-flight key constexpr int DP4 = D4 / NL; // float4s per lane per key row - constexpr int GP4 = 64 / 4; // packed words per quant group (group_size 64) using T4 = metal::vec; threadgroup T4 sK[C * D4];