From 727c58e6d111f821c3e06332ed0fcb5ac16b7419 Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Fri, 4 Sep 2026 01:10:12 +0200 Subject: [PATCH 1/4] added global scale for qmm --- mlx/backend/cuda/quantized/quantized.cpp | 7 +- mlx/backend/metal/jit_kernels.cpp | 84 ++++++--- mlx/backend/metal/kernels.h | 6 +- mlx/backend/metal/kernels/fp_quantized.h | 139 ++++++++++---- mlx/backend/metal/kernels/fp_quantized.metal | 30 ++- mlx/backend/metal/kernels/fp_quantized_nax.h | 107 ++++++++--- .../metal/kernels/fp_quantized_nax.metal | 24 ++- mlx/backend/metal/nojit_kernels.cpp | 2 + mlx/backend/metal/quantized.cpp | 175 +++++++++++++----- mlx/ops.cpp | 43 ++++- mlx/ops.h | 1 + mlx/primitives.cpp | 4 + python/src/ops.cpp | 6 +- python/tests/test_quantized.py | 77 ++++++++ 14 files changed, 560 insertions(+), 145 deletions(-) diff --git a/mlx/backend/cuda/quantized/quantized.cpp b/mlx/backend/cuda/quantized/quantized.cpp index 4d25f3c3e0..03e6b4f028 100644 --- a/mlx/backend/cuda/quantized/quantized.cpp +++ b/mlx/backend/cuda/quantized/quantized.cpp @@ -156,11 +156,16 @@ void GatherQMM::eval_gpu(const std::vector& inputs, array& out) { auto& s = stream(); auto& encoder = cu::get_command_encoder(s); + if (mode_ != QuantizationMode::Affine && inputs.size() == 6) { + throw std::runtime_error( + "[GatherQMM] Global scale is only supported on the Metal backend."); + } + array x = ensure_row_contiguous(inputs[0], encoder, s); const array& w = inputs[1]; const array& scales = inputs[2]; std::optional biases; - if (inputs.size() == 6) { + if (mode_ == QuantizationMode::Affine) { biases = inputs[3]; } array lhs_indices = diff --git a/mlx/backend/metal/jit_kernels.cpp b/mlx/backend/metal/jit_kernels.cpp index 8dfe30a15c..807a6550b7 100644 --- a/mlx/backend/metal/jit_kernels.cpp +++ b/mlx/backend/metal/jit_kernels.cpp @@ -1054,28 +1054,44 @@ MTL::ComputePipelineState* get_gather_qmm_kernel( int bk, int wm, int wn, - bool transpose) { + bool transpose, + bool has_global_scale) { const auto& lib_name = kernel_name; auto lib = d.get_library(lib_name, [&]() { std::string kernel_source; concatenate( kernel_source, metal::utils(), metal::quantized_utils(), metal::gemm()); bool is_affine = mode == "affine"; + // Only the fp kernel takes a global scale. + auto template_def = is_affine ? get_template_definition( + lib_name, + "affine_gather_qmm_rhs", + get_type_string(x.dtype()), + group_size, + bits, + bm, + bn, + bk, + wm, + wn, + transpose) + : get_template_definition( + lib_name, + "fp_gather_qmm_rhs", + get_type_string(x.dtype()), + group_size, + bits, + bm, + bn, + bk, + wm, + wn, + transpose, + has_global_scale); concatenate( kernel_source, is_affine ? metal::quantized() : metal::fp_quantized(), - get_template_definition( - lib_name, - (is_affine ? "affine" : "fp") + std::string("_gather_qmm_rhs"), - get_type_string(x.dtype()), - group_size, - bits, - bm, - bn, - bk, - wm, - wn, - transpose)); + template_def); return kernel_source; }); return d.get_kernel(kernel_name, lib, hash_name, func_consts); @@ -1255,7 +1271,8 @@ MTL::ComputePipelineState* get_gather_qmm_nax_kernel( int bk, int wm, int wn, - bool transpose) { + bool transpose, + bool has_global_scale) { const auto& lib_name = kernel_name; auto lib = d.get_library(lib_name, [&]() { std::string kernel_source; @@ -1265,21 +1282,36 @@ MTL::ComputePipelineState* get_gather_qmm_nax_kernel( metal::gemm_nax(), metal::quantized_utils()); bool is_affine = mode == "affine"; + // Only the fp kernel takes a global scale. + auto template_def = is_affine ? get_template_definition( + lib_name, + "affine_gather_qmm_rhs_nax", + get_type_string(x.dtype()), + group_size, + bits, + bm, + bn, + bk, + wm, + wn, + transpose) + : get_template_definition( + lib_name, + "fp_gather_qmm_rhs_nax", + get_type_string(x.dtype()), + group_size, + bits, + bm, + bn, + bk, + wm, + wn, + transpose, + has_global_scale); concatenate( kernel_source, is_affine ? metal::quantized_nax() : metal::fp_quantized_nax(), - get_template_definition( - lib_name, - (is_affine ? "affine" : "fp") + std::string("_gather_qmm_rhs_nax"), - get_type_string(x.dtype()), - group_size, - bits, - bm, - bn, - bk, - wm, - wn, - transpose)); + template_def); return kernel_source; }); return d.get_kernel(kernel_name, lib, hash_name, func_consts); diff --git a/mlx/backend/metal/kernels.h b/mlx/backend/metal/kernels.h index 18a56ebf1b..42888f0a78 100644 --- a/mlx/backend/metal/kernels.h +++ b/mlx/backend/metal/kernels.h @@ -321,7 +321,8 @@ MTL::ComputePipelineState* get_gather_qmm_kernel( int bk, int wm, int wn, - bool transpose); + bool transpose, + bool has_global_scale); MTL::ComputePipelineState* get_steel_gemm_fused_nax_kernel( metal::Device& d, @@ -400,7 +401,8 @@ MTL::ComputePipelineState* get_gather_qmm_nax_kernel( int bk, int wm, int wn, - bool transpose); + bool transpose, + bool has_global_scale); MTL::ComputePipelineState* get_steel_attention_kernel( metal::Device& d, diff --git a/mlx/backend/metal/kernels/fp_quantized.h b/mlx/backend/metal/kernels/fp_quantized.h index 8c963030f2..2495001f0a 100644 --- a/mlx/backend/metal/kernels/fp_quantized.h +++ b/mlx/backend/metal/kernels/fp_quantized.h @@ -136,13 +136,12 @@ inline void qouter(const thread uint8_t* w, U x, U scale, thread U* result) { } template -inline void dequantize(uint8_t w, U scale, threadgroup U* w_local) { - const float s = float(scale); +inline void dequantize(uint8_t w, float scale, threadgroup U* w_local) { if constexpr (bits == 4) { - w_local[0] = static_cast(s * Dequantize<4, float>{}(w)); - w_local[1] = static_cast(s * Dequantize<4, float>{}(w >> 4)); + w_local[0] = static_cast(scale * Dequantize<4, float>{}(w)); + w_local[1] = static_cast(scale * Dequantize<4, float>{}(w >> 4)); } else { - w_local[0] = static_cast(s * Dequantize<8, float>{}(w)); + w_local[0] = static_cast(scale * Dequantize<8, float>{}(w)); } } @@ -154,7 +153,8 @@ template < short reduction_dim, short tgp_size, short group_size, - short bits> + short bits, + bool has_global_scale = false> struct QuantizedBlockLoader { MLX_MTL_CONST short pack_factor = get_pack_factor<8, bits>(); MLX_MTL_CONST short bytes_per_pack = get_bytes_per_pack(); @@ -180,6 +180,9 @@ struct QuantizedBlockLoader { threadgroup T* dst; const device uint8_t* src; const device uint8_t* scales; + // nvfp4 tensor scale, folded into the group scale as fp_dequantize does. + // Kept in float: it is ~1e-5, so in fp16 small scales lose most bits. + float inv_scale_enc = 1.0f; QuantizedBlockLoader( const device uint8_t* src_, @@ -187,7 +190,8 @@ struct QuantizedBlockLoader { const int src_ld_, threadgroup T* dst_, ushort simd_group_id [[simdgroup_index_in_threadgroup]], - ushort simd_lane_id [[thread_index_in_simdgroup]]) thread + ushort simd_lane_id [[thread_index_in_simdgroup]], + const device float* global_scale = nullptr) thread : src_ld(src_ld_), tile_stride( reduction_dim ? BCOLS_PACKED* bytes_per_pack @@ -202,14 +206,19 @@ struct QuantizedBlockLoader { bj * bytes_per_pack), scales( scales_ + bi * src_ld / group_size + - (bj * pack_factor) / group_size) {} + (bj * pack_factor) / group_size) { + if constexpr (has_global_scale) { + inv_scale_enc = *global_scale / (F8E4M3_MAX * F4E2M1_MAX); + } + } void load_unsafe() const thread { if (BCOLS_PACKED * BROWS < tgp_size && bi >= BROWS) { return; } - T scale = dequantize_scale(*scales); + float scale = + float(dequantize_scale(*scales)) * inv_scale_enc; for (int i = 0; i < n_reads; i++) { dequantize( src[i * bytes_per_pack], scale, dst + i * pack_factor); @@ -235,7 +244,8 @@ struct QuantizedBlockLoader { return; } - T scale = dequantize_scale(*scales); + float scale = + float(dequantize_scale(*scales)) * inv_scale_enc; for (int i = 0; i < n_reads; i++) { dequantize( src[i * bytes_per_pack], scale, dst + i * pack_factor); @@ -788,12 +798,14 @@ template < const int group_size, const int bits, const bool aligned_N, + const bool has_global_scale = false, const int BM = 32, const int BK = 32, const int BN = 32> METAL_FUNC void fp_qmm_t_impl( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, device T* y, threadgroup T* Xs, @@ -831,7 +843,8 @@ METAL_FUNC void fp_qmm_t_impl( 1, WM * WN * SIMD_SIZE, group_size, - bits>; + bits, + has_global_scale>; // Set the block const int K_w = K * bytes_per_pack / pack_factor; @@ -850,7 +863,7 @@ METAL_FUNC void fp_qmm_t_impl( const short num_els = min(BM, M - y_row); const short num_outs = min(BN, N - y_col); loader_x_t loader_x(x, K, Xs, simd_gid, simd_lid); - loader_w_t loader_w(wl, scales, K, Ws, simd_gid, simd_lid); + loader_w_t loader_w(wl, scales, K, Ws, simd_gid, simd_lid, global_scale); mma_t mma_op(simd_gid, simd_lid); if (num_els < BM) { @@ -913,12 +926,14 @@ template < typename T, int group_size, int bits, + bool has_global_scale = false, int BM = 32, int BK = 32, int BN = 32> METAL_FUNC void fp_qmm_n_impl( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, device T* y, threadgroup T* Xs, @@ -956,7 +971,8 @@ METAL_FUNC void fp_qmm_n_impl( 0, WM * WN * SIMD_SIZE, group_size, - bits>; + bits, + has_global_scale>; auto wl = (const device uint8_t*)w; @@ -971,7 +987,7 @@ METAL_FUNC void fp_qmm_n_impl( // Make the x loader and mma operation const short num_els = min(BM, M - y_row); loader_x_t loader_x(x, K, Xs, simd_gid, simd_lid); - loader_w_t loader_w(wl, scales, N, Ws, simd_gid, simd_lid); + loader_w_t loader_w(wl, scales, N, Ws, simd_gid, simd_lid, global_scale); mma_t mma_op(simd_gid, simd_lid); if (num_els < BM) { @@ -1078,11 +1094,12 @@ METAL_FUNC void adjust_matrix_offsets( y += tid.z * output_stride; } -template +template METAL_FUNC void adjust_matrix_offsets( const device T*& x, const device uint32_t*& w, const device uint8_t*& scales, + const device float*& global_scale, const device uint32_t* lhs_indices, const device uint32_t* rhs_indices, device T*& y, @@ -1125,6 +1142,10 @@ METAL_FUNC void adjust_matrix_offsets( w += idx.x; scales += idx.y; } + // One global scale per expert, contiguous over the batch dims of w. + if constexpr (has_global_scale) { + global_scale += w_idx; + } y += tid.z * output_stride; } @@ -1488,8 +1509,22 @@ template < s_strides, tid); } - fp_qmm_t_impl( - w, scales, x, y, Xs, Ws, K, N, M, K, tid, lid, simd_gid, simd_lid); + fp_qmm_t_impl( + w, + scales, + nullptr, + x, + y, + Xs, + Ws, + K, + N, + M, + K, + tid, + lid, + simd_gid, + simd_lid); } template < @@ -1545,8 +1580,8 @@ template < tid); } - fp_qmm_n_impl( - w, scales, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); + fp_qmm_n_impl( + w, scales, nullptr, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); } template @@ -1575,10 +1610,11 @@ template uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { int M = x_shape[x_batch_ndims]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -1634,10 +1670,11 @@ template uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { int M = x_shape[x_batch_ndims]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -1693,10 +1730,11 @@ template uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { int M = x_shape[x_batch_ndims]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -1732,12 +1770,14 @@ template < const int group_size, const int bits, const bool aligned_N, + const bool has_global_scale = false, const int BM = 32, const int BK = 32, const int BN = 32> [[kernel]] void fp_gather_qmm_t( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, const device uint32_t* lhs_indices, const device uint32_t* rhs_indices, @@ -1767,10 +1807,11 @@ template < threadgroup T Xs[BM * BK_padded]; threadgroup T Ws[BN * BK_padded]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -1787,8 +1828,22 @@ template < w_strides, s_strides, tid); - fp_qmm_t_impl( - w, scales, x, y, Xs, Ws, K, N, M, K, tid, lid, simd_gid, simd_lid); + fp_qmm_t_impl( + w, + scales, + global_scale, + x, + y, + Xs, + Ws, + K, + N, + M, + K, + tid, + lid, + simd_gid, + simd_lid); } template < @@ -1828,9 +1883,10 @@ template < scales += k_start / group_size; y += tid.z * static_cast(split_k_partition_stride); - fp_qmm_t_impl( + fp_qmm_t_impl( (const device uint32_t*)wl, scales, + nullptr, x, y, Xs, @@ -1856,6 +1912,7 @@ template < [[kernel]] void fp_gather_qmm_n( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, const device uint32_t* lhs_indices, const device uint32_t* rhs_indices, @@ -1886,10 +1943,11 @@ template < threadgroup T Xs[BM * BK_padded]; threadgroup T Ws[BK * BN_padded]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -1906,8 +1964,21 @@ template < w_strides, s_strides, tid); - fp_qmm_n_impl( - w, scales, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); + fp_qmm_n_impl( + w, + scales, + global_scale, + x, + y, + Xs, + Ws, + K, + N, + M, + tid, + lid, + simd_gid, + simd_lid); } template < @@ -1919,11 +1990,13 @@ template < int BK, int WM, int WN, - bool transpose> + bool transpose, + bool has_global_scale = false> [[kernel]] void fp_gather_qmm_rhs( const device T* x, const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device uint32_t* indices, device T* y, const constant int& M, @@ -1959,7 +2032,8 @@ template < transpose, WM * WN * SIMD_SIZE, group_size, - bits>; + bits, + has_global_scale>; threadgroup T Xs[BM * BK_padded]; threadgroup T Ws[transpose ? BN * BK_padded : BK * BN_padded]; @@ -2025,7 +2099,8 @@ template < transpose ? K : N, Ws, simd_group_id, - simd_lane_id); + simd_lane_id, + global_scale + index); // Matrices are all aligned check nothing if (align_M && align_N) { diff --git a/mlx/backend/metal/kernels/fp_quantized.metal b/mlx/backend/metal/kernels/fp_quantized.metal index d8d462288a..8cf4639f65 100644 --- a/mlx/backend/metal/kernels/fp_quantized.metal +++ b/mlx/backend/metal/kernels/fp_quantized.metal @@ -47,6 +47,17 @@ bits, \ aligned) +#define instantiate_quantized_aligned_hgs(mode, name, type, aligned, group_size, bits) \ + instantiate_quantized_aligned(mode, name, type, aligned, group_size, bits) \ + instantiate_kernel( \ + #mode "_" #name "_" #type "_gs_" #group_size "_b_" #bits "_alN_" #aligned "_hgs", \ + fp_ ## name, \ + type, \ + group_size, \ + bits, \ + aligned, \ + true) + #define instantiate_quantized_aligned_batched(mode, name, type, aligned, batched, group_size, bits) \ instantiate_kernel( \ #mode "_" #name "_" #type "_gs_" #group_size "_b_" #bits "_alN_" #aligned "_batch_" #batched, \ @@ -123,7 +134,20 @@ bk, \ wm, \ wn, \ - transpose) + transpose) \ + instantiate_kernel( \ + #mode "_" #name "_" #type "_gs_" #group_size "_b_" #bits "_bm_" #bm "_bn_" #bn "_bk_" #bk "_wm_" #wm "_wn_" #wn "_hgs", \ + func, \ + type, \ + group_size, \ + bits, \ + bm, \ + bn, \ + bk, \ + wm, \ + wn, \ + transpose, \ + true) #define instantiate_quantized_batched_wrap(name, type, mode, group_size, bits) \ instantiate_quantized_batched(mode, name, type, 1, group_size, bits) \ @@ -142,8 +166,8 @@ instantiate_quantized(mode, gather_qmm_n, type, group_size, bits) #define instantiate_quantized_all_aligned(type, mode, group_size, bits) \ - instantiate_quantized_aligned(mode, gather_qmm_t, type, true, group_size, bits) \ - instantiate_quantized_aligned(mode, gather_qmm_t, type, false, group_size, bits) \ + instantiate_quantized_aligned_hgs(mode, gather_qmm_t, type, true, group_size, bits) \ + instantiate_quantized_aligned_hgs(mode, gather_qmm_t, type, false, group_size, bits) \ instantiate_quantized_aligned_batched(mode, qmm_t, type, true, 1, group_size, bits) \ instantiate_quantized_aligned_batched(mode, qmm_t, type, true, 0, group_size, bits) \ instantiate_quantized_aligned_batched(mode, qmm_t, type, false, 1, group_size, bits) \ diff --git a/mlx/backend/metal/kernels/fp_quantized_nax.h b/mlx/backend/metal/kernels/fp_quantized_nax.h index 946bce7868..2a62b587b0 100644 --- a/mlx/backend/metal/kernels/fp_quantized_nax.h +++ b/mlx/backend/metal/kernels/fp_quantized_nax.h @@ -60,13 +60,12 @@ struct Dequantize { }; template -inline void dequantize(uint8_t w, U scale, threadgroup U* w_local) { - const float s = float(scale); +inline void dequantize(uint8_t w, float scale, threadgroup U* w_local) { if constexpr (bits == 4) { - w_local[0] = static_cast(s * Dequantize<4, float>{}(w)); - w_local[1] = static_cast(s * Dequantize<4, float>{}(w >> 4)); + w_local[0] = static_cast(scale * Dequantize<4, float>{}(w)); + w_local[1] = static_cast(scale * Dequantize<4, float>{}(w >> 4)); } else { - w_local[0] = static_cast(s * Dequantize<8, float>{}(w)); + w_local[0] = static_cast(scale * Dequantize<8, float>{}(w)); } } @@ -78,7 +77,8 @@ template < short reduction_dim, short tgp_size, short group_size, - short bits> + short bits, + bool has_global_scale = false> struct QuantizedBlockLoader { MLX_MTL_CONST short pack_factor = get_pack_factor<8, bits>(); MLX_MTL_CONST short bytes_per_pack = get_bytes_per_pack(); @@ -106,6 +106,9 @@ struct QuantizedBlockLoader { threadgroup T* dst; const device uint8_t* src; const device uint8_t* scales; + // nvfp4 tensor scale, folded into the group scale as fp_dequantize does. + // Kept in float: it is ~1e-5, so in fp16 small scales lose most bits. + float inv_scale_enc = 1.0f; QuantizedBlockLoader( const device uint8_t* src_, @@ -113,7 +116,8 @@ struct QuantizedBlockLoader { const int src_ld_, threadgroup T* dst_, ushort simd_group_id [[simdgroup_index_in_threadgroup]], - ushort simd_lane_id [[thread_index_in_simdgroup]]) thread + ushort simd_lane_id [[thread_index_in_simdgroup]], + const device float* global_scale = nullptr) thread : src_ld(src_ld_), tile_stride( reduction_dim ? BCOLS_PACKED* bytes_per_pack @@ -126,7 +130,11 @@ struct QuantizedBlockLoader { dst(dst_ + bi * dst_ld + bj * pack_factor), src(src_ + bi * src_ld * bytes_per_pack / pack_factor + bj * bytes_per_pack), - scales(scales_ + bi * src_ld / group_size + group_id) {} + scales(scales_ + bi * src_ld / group_size + group_id) { + if constexpr (has_global_scale) { + inv_scale_enc = *global_scale / (F8E4M3_MAX * F4E2M1_MAX); + } + } void load_unsafe() const thread { if (BCOLS_PACKED * BROWS < tgp_size && bi >= BROWS) { @@ -135,7 +143,8 @@ struct QuantizedBlockLoader { int k = 0; for (int i = 0; i < n_steps_per_read; i++) { - T scale = dequantize_scale(scales[i]); + float scale = + float(dequantize_scale(scales[i])) * inv_scale_enc; for (int j = 0; j < n_reads_per_scale; j++) { dequantize( src[k * bytes_per_pack], scale, dst + k * pack_factor); @@ -165,7 +174,8 @@ struct QuantizedBlockLoader { int k = 0; for (int i = 0; i < n_steps_per_read; i++) { - T scale = dequantize_scale(scales[i]); + float scale = + float(dequantize_scale(scales[i])) * inv_scale_enc; for (int j = 0; j < n_reads_per_scale; j++) { dequantize( src[k * bytes_per_pack], scale, dst + k * pack_factor); @@ -191,6 +201,7 @@ template < const int group_size, const int bits, const bool aligned_N, + const bool has_global_scale = false, const int BM = 64, const int BK = 64, const int BN = 64, @@ -200,6 +211,7 @@ template < METAL_FUNC void fp_qmm_t_impl( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, device T* y, threadgroup Wtype* Ws, @@ -229,7 +241,8 @@ METAL_FUNC void fp_qmm_t_impl( 1, WM * WN * SIMD_SIZE, group_size, - bits>; + bits, + has_global_scale>; // Set the block const int K_w = K * bytes_per_pack / pack_factor; @@ -245,7 +258,7 @@ METAL_FUNC void fp_qmm_t_impl( y += y_row * static_cast(N) + y_col; // Make the weight loader - loader_w_t loader_w(wl, scales, K, Ws, simd_gid, simd_lid); + loader_w_t loader_w(wl, scales, K, Ws, simd_gid, simd_lid, global_scale); constexpr short SM = BM / WM; constexpr short SN = BN / WN; @@ -335,6 +348,7 @@ template < typename T, const int group_size, const int bits, + const bool has_global_scale = false, const int BM = 64, const int BK = 64, const int BN = 64, @@ -344,6 +358,7 @@ template < METAL_FUNC void fp_qmm_n_impl( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, device T* y, threadgroup T* Ws, @@ -373,7 +388,8 @@ METAL_FUNC void fp_qmm_n_impl( 0, WM * WN * SIMD_SIZE, group_size, - bits>; + bits, + has_global_scale>; // Set the block const int K_w = K * bytes_per_pack / pack_factor; @@ -391,7 +407,7 @@ METAL_FUNC void fp_qmm_n_impl( // Make the x loader and mma operation // const short num_els = min(BM, M - y_row); // const short num_outs = min(BN, N - y_col); - loader_w_t loader_w(wl, scales, K, Ws, simd_gid, simd_lid); + loader_w_t loader_w(wl, scales, K, Ws, simd_gid, simd_lid, global_scale); constexpr short SM = BM / WM; constexpr short SN = BN / WN; @@ -486,11 +502,12 @@ METAL_FUNC void adjust_matrix_offsets( y += tid.z * output_stride; } -template +template METAL_FUNC void adjust_matrix_offsets( const device T*& x, const device uint32_t*& w, const device S*& scales, + const device float*& global_scale, const device uint32_t* lhs_indices, const device uint32_t* rhs_indices, device T*& y, @@ -533,6 +550,10 @@ METAL_FUNC void adjust_matrix_offsets( w += idx.x; scales += idx.y; } + // One global scale per expert, contiguous over the batch dims of w. + if constexpr (has_global_scale) { + global_scale += w_idx; + } y += tid.z * output_stride; } @@ -589,8 +610,19 @@ template < s_strides, tid); } - fp_qmm_t_impl( - w, scales, x, y, Ws, K, N, M, tid, lid, simd_gid, simd_lid); + fp_qmm_t_impl< + T, + group_size, + bits, + aligned_N, + false, + BM, + BK, + BN, + WM, + WN, + Wtype>( + w, scales, nullptr, x, y, Ws, K, N, M, tid, lid, simd_gid, simd_lid); } template < @@ -648,8 +680,8 @@ template < tid); } - fp_qmm_n_impl( - w, scales, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); + fp_qmm_n_impl( + w, scales, nullptr, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); } template < @@ -662,10 +694,12 @@ template < const int BN = 64, const int WM = 2, const int WN = 2, + const bool has_global_scale = false, typename Wtype = bfloat> [[kernel]] void fp_gather_qmm_t_nax( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, const device uint32_t* lhs_indices, const device uint32_t* rhs_indices, @@ -694,10 +728,11 @@ template < threadgroup Wtype Ws[BN * BK_padded]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -714,8 +749,19 @@ template < w_strides, s_strides, tid); - fp_qmm_t_impl( - w, scales, x, y, Ws, K, N, M, tid, lid, simd_gid, simd_lid); + fp_qmm_t_impl< + T, + group_size, + bits, + aligned_N, + has_global_scale, + BM, + BK, + BN, + WM, + WN, + Wtype>( + w, scales, global_scale, x, y, Ws, K, N, M, tid, lid, simd_gid, simd_lid); } template < @@ -727,10 +773,12 @@ template < const int BN = 64, const int WM = 2, const int WN = 2, + const bool has_global_scale = false, typename Wtype = bfloat> [[kernel]] void fp_gather_qmm_n_nax( const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device T* x, const device uint32_t* lhs_indices, const device uint32_t* rhs_indices, @@ -761,10 +809,11 @@ template < threadgroup T Xs[BM * BK_padded]; threadgroup T Ws[BK * BN_padded]; - adjust_matrix_offsets( + adjust_matrix_offsets( x, w, scales, + global_scale, lhs_indices, rhs_indices, y, @@ -781,8 +830,8 @@ template < w_strides, s_strides, tid); - fp_qmm_n_impl( - w, scales, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); + fp_qmm_n_impl( + w, scales, nullptr, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); } template < @@ -795,11 +844,13 @@ template < int WM, int WN, bool transpose, + bool has_global_scale = false, typename Wtype = bfloat> [[kernel]] void fp_gather_qmm_rhs_nax( const device T* x, const device uint32_t* w, const device uint8_t* scales, + const device float* global_scale, const device uint32_t* indices, device T* y, const constant int& M, @@ -821,7 +872,8 @@ template < transpose, WM * WN * SIMD_SIZE, group_size, - bits>; + bits, + has_global_scale>; threadgroup Wtype Ws[transpose ? BN * BK_padded : BK * BN_padded]; @@ -913,7 +965,8 @@ template < transpose ? K : N, Ws, simd_group_id, - simd_lane_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) { diff --git a/mlx/backend/metal/kernels/fp_quantized_nax.metal b/mlx/backend/metal/kernels/fp_quantized_nax.metal index 771b2a963a..32cb82bb4e 100644 --- a/mlx/backend/metal/kernels/fp_quantized_nax.metal +++ b/mlx/backend/metal/kernels/fp_quantized_nax.metal @@ -24,7 +24,14 @@ type, \ group_size, \ bits, \ - aligned, bm, bk, bn, wm, wn) + aligned, bm, bk, bn, wm, wn) \ + instantiate_kernel( \ + #mode "_" #name "_" #type "_gs_" #group_size "_b_" #bits "_bm" #bm "_bn" #bn "_bk" #bk "_wm" #wm "_wn" #wn "_alN_" #aligned "_hgs", \ + fp_ ## name, \ + type, \ + group_size, \ + bits, \ + aligned, bm, bk, bn, wm, wn, true) #define instantiate_quantized_aligned_batched(mode, name, type, bm, bn, bk, wm, wn, aligned, batched, group_size, bits) \ instantiate_kernel( \ @@ -48,7 +55,20 @@ bk, \ wm, \ wn, \ - transpose) + transpose) \ + instantiate_kernel( \ + #mode "_" #name "_" #type "_gs_" #group_size "_b_" #bits "_bm_" #bm "_bn_" #bn "_bk_" #bk "_wm_" #wm "_wn_" #wn "_hgs", \ + func, \ + type, \ + group_size, \ + bits, \ + bm, \ + bn, \ + bk, \ + wm, \ + wn, \ + transpose, \ + true) #define instantiate_quantized_all_aligned(type, mode, group_size, bits) \ diff --git a/mlx/backend/metal/nojit_kernels.cpp b/mlx/backend/metal/nojit_kernels.cpp index 3795a6fb22..f1141e9792 100644 --- a/mlx/backend/metal/nojit_kernels.cpp +++ b/mlx/backend/metal/nojit_kernels.cpp @@ -382,6 +382,7 @@ MTL::ComputePipelineState* get_gather_qmm_kernel( int, int, int, + bool, bool) { return d.get_kernel(kernel_name, hash_name, func_consts); } @@ -473,6 +474,7 @@ MTL::ComputePipelineState* get_gather_qmm_nax_kernel( int, int, int, + bool, bool) { return d.get_kernel(kernel_name, hash_name, func_consts); } diff --git a/mlx/backend/metal/quantized.cpp b/mlx/backend/metal/quantized.cpp index 013bc6d333..82c727d238 100644 --- a/mlx/backend/metal/quantized.cpp +++ b/mlx/backend/metal/quantized.cpp @@ -920,6 +920,7 @@ void gather_qmm_nax( const array& w, const array& scales, const std::optional& biases, + const std::optional& global_scale, const array& lhs_indices, const array& rhs_indices, array& out, @@ -966,23 +967,39 @@ void gather_qmm_nax( wm, "_wn", wn, - transpose ? (aligned ? "_alN_true" : "_alN_false") : ""); + transpose ? (aligned ? "_alN_true" : "_alN_false") : "", + global_scale ? "_hgs" : ""); MTL::ComputePipelineState* kernel; if (transpose) { - kernel = get_qmm_nax_kernel_wrapped( - d, - kname, - "gather_qmm_t_nax_", - mode, - type_string, - group_size, - bits, - aligned, - bm, - bk, - bn, - wm, - wn); + kernel = global_scale ? get_qmm_nax_kernel_wrapped( + d, + kname, + "gather_qmm_t_nax_", + mode, + type_string, + group_size, + bits, + aligned, + bm, + bk, + bn, + wm, + wn, + true) + : get_qmm_nax_kernel_wrapped( + d, + kname, + "gather_qmm_t_nax_", + mode, + type_string, + group_size, + bits, + aligned, + bm, + bk, + bn, + wm, + wn); } else { kernel = get_qmm_nax_kernel_wrapped( d, @@ -1002,12 +1019,14 @@ void gather_qmm_nax( auto& compute_encoder = metal::get_command_encoder(s); compute_encoder.set_compute_pipeline_state(kernel); - int c = 0; - compute_encoder.set_input_array(w, c++); - compute_encoder.set_input_array(scales, c++); + compute_encoder.set_input_array(w, 0); + compute_encoder.set_input_array(scales, 1); if (biases) { - compute_encoder.set_input_array(*biases, c++); + compute_encoder.set_input_array(*biases, 2); + } else if (global_scale) { + compute_encoder.set_input_array(*global_scale, 2); } + int c = 3; compute_encoder.set_input_array(x, c++); compute_encoder.set_input_array(lhs_indices, c++); compute_encoder.set_input_array(rhs_indices, c++); @@ -1221,6 +1240,7 @@ void gather_qmm( const array& w, const array& scales, const std::optional& biases, + const std::optional& global_scale, const array& lhs_indices, const array& rhs_indices, array& out, @@ -1240,6 +1260,7 @@ void gather_qmm( /* const array& w = */ w, /* const array& scales = */ scales, /* const std::optional& biases = */ biases, + /* const std::optional& global_scale = */ global_scale, /* const array& lhs_indices = */ lhs_indices, /* const array& rhs_indices = */ rhs_indices, /* array& out = */ out, @@ -1275,25 +1296,55 @@ void gather_qmm( group_size, "_b_", bits, - transpose ? (aligned ? "_alN_true" : "_alN_false") : ""); + transpose ? (aligned ? "_alN_true" : "_alN_false") : "", + global_scale ? "_hgs" : ""); MTL::ComputePipelineState* kernel; if (transpose) { - kernel = get_quantized_kernel_wrapped( - d, kname, "gather_qmm_t", mode, type_string, group_size, bits, aligned); + kernel = global_scale ? get_quantized_kernel_wrapped( + d, + kname, + "gather_qmm_t", + mode, + type_string, + group_size, + bits, + aligned, + true) + : get_quantized_kernel_wrapped( + d, + kname, + "gather_qmm_t", + mode, + type_string, + group_size, + bits, + aligned); } else { - kernel = get_quantized_kernel_wrapped( - d, kname, "gather_qmm_n", mode, type_string, group_size, bits); + kernel = global_scale + ? get_quantized_kernel_wrapped( + d, + kname, + "gather_qmm_n", + mode, + type_string, + group_size, + bits, + true) + : get_quantized_kernel_wrapped( + d, kname, "gather_qmm_n", mode, type_string, group_size, bits); } auto& compute_encoder = metal::get_command_encoder(s); compute_encoder.set_compute_pipeline_state(kernel); - int c = 0; - compute_encoder.set_input_array(w, c++); - compute_encoder.set_input_array(scales, c++); + compute_encoder.set_input_array(w, 0); + compute_encoder.set_input_array(scales, 1); if (biases) { - compute_encoder.set_input_array(*biases, c++); + compute_encoder.set_input_array(*biases, 2); + } else if (global_scale) { + compute_encoder.set_input_array(*global_scale, 2); } + int c = 3; compute_encoder.set_input_array(x, c++); compute_encoder.set_input_array(lhs_indices, c++); compute_encoder.set_input_array(rhs_indices, c++); @@ -1452,6 +1503,7 @@ void gather_qmm_rhs_nax( const array& w_, const array& scales_, const std::optional& biases_, + const std::optional& global_scale, const array& indices_, array& out, bool transpose, @@ -1491,6 +1543,10 @@ void gather_qmm_rhs_nax( if (biases_) { biases = ensure_row_contiguous(*biases_, d, s); } + std::optional gs; + if (global_scale) { + gs = ensure_row_contiguous(*global_scale, d, s); + } // Use smaller bm for many experts and few tokens. int E = w.size() / w.shape(-1) / w.shape(-2); @@ -1524,7 +1580,8 @@ void gather_qmm_rhs_nax( "_wm_", wm, "_wn_", - wn); + wn, + global_scale ? "_hgs" : ""); metal::MTLFCList func_consts = { {&align_M, MTL::DataType::DataTypeBool, 200}, @@ -1561,19 +1618,22 @@ void gather_qmm_rhs_nax( bk, wm, wn, - transpose); + transpose, + global_scale.has_value()); 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); - int c = 0; - compute_encoder.set_input_array(x, c++); - compute_encoder.set_input_array(w, c++); - compute_encoder.set_input_array(scales, c++); + compute_encoder.set_input_array(x, 0); + compute_encoder.set_input_array(w, 1); + compute_encoder.set_input_array(scales, 2); if (biases) { - compute_encoder.set_input_array(*biases, c++); + compute_encoder.set_input_array(*biases, 3); + } else if (gs) { + compute_encoder.set_input_array(*gs, 3); } + int c = 4; compute_encoder.set_input_array(indices, c++); compute_encoder.set_output_array(out, c++); compute_encoder.set_bytes(M, c++); @@ -1588,6 +1648,7 @@ void gather_qmm_rhs( const array& w_, const array& scales_, const std::optional& biases_, + const std::optional& global_scale, const array& indices_, array& out, bool transpose, @@ -1606,6 +1667,7 @@ void gather_qmm_rhs( /* const array& w_ = */ w_, /* const array& scales_ = */ scales_, /* const std::optional& biases_ = */ biases_, + /* const std::optional& global_scale = */ global_scale, /* const array& indices_ = */ indices_, /* array& out = */ out, /* bool transpose = */ transpose, @@ -1647,6 +1709,10 @@ void gather_qmm_rhs( if (biases_) { biases = ensure_row_contiguous(*biases_, d, s); } + std::optional gs; + if (global_scale) { + gs = ensure_row_contiguous(*global_scale, d, s); + } // TODO: Tune the block sizes int bm = 16, bn = 32, bk = 32; @@ -1677,7 +1743,8 @@ void gather_qmm_rhs( "_wm_", wm, "_wn_", - wn); + wn, + global_scale ? "_hgs" : ""); metal::MTLFCList func_consts = { {&align_M, MTL::DataType::DataTypeBool, 200}, @@ -1714,19 +1781,22 @@ void gather_qmm_rhs( bk, wm, wn, - transpose); + transpose, + global_scale.has_value()); 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); - int c = 0; - compute_encoder.set_input_array(x, c++); - compute_encoder.set_input_array(w, c++); - compute_encoder.set_input_array(scales, c++); + compute_encoder.set_input_array(x, 0); + compute_encoder.set_input_array(w, 1); + compute_encoder.set_input_array(scales, 2); if (biases) { - compute_encoder.set_input_array(*biases, c++); + compute_encoder.set_input_array(*biases, 3); + } else if (gs) { + compute_encoder.set_input_array(*gs, 3); } + int c = 4; compute_encoder.set_input_array(indices, c++); compute_encoder.set_output_array(out, c++); compute_encoder.set_bytes(M, c++); @@ -1883,9 +1953,13 @@ void GatherQMM::eval_gpu(const std::vector& inputs, array& out) { array x = ensure_row_contiguous_matrix(inputs[0], d, s); array w = ensure_row_contiguous_matrix(inputs[1], d, s); array scales = ensure_row_contiguous_matrix(inputs[2], d, s); + // Affine gets biases at index 3, nvfp4 an optional global scale. std::optional biases = std::nullopt; - if (inputs.size() == 6) { + std::optional global_scale = std::nullopt; + if (mode_ == QuantizationMode::Affine) { biases = ensure_row_contiguous_matrix(inputs[3], d, s); + } else if (inputs.size() == 6) { + global_scale = ensure_row_contiguous(inputs[3], d, s); } const array& lhs_indices = inputs[inputs.size() - 2]; const array& rhs_indices = inputs[inputs.size() - 1]; @@ -1908,6 +1982,7 @@ void GatherQMM::eval_gpu(const std::vector& inputs, array& out) { w, scales, biases, + global_scale, rhs_indices, out, transpose_, @@ -1929,6 +2004,7 @@ void GatherQMM::eval_gpu(const std::vector& inputs, array& out) { w, scales, biases, + global_scale, lhs_indices, rhs_indices, out, @@ -1950,7 +2026,7 @@ void GatherQMM::eval_gpu(const std::vector& inputs, array& out) { w, scales, biases, - std::nullopt, + global_scale, lhs_indices, rhs_indices, out, @@ -1970,7 +2046,7 @@ void GatherQMM::eval_gpu(const std::vector& inputs, array& out) { w, scales, biases, - std::nullopt, + global_scale, lhs_indices, rhs_indices, out, @@ -2087,6 +2163,15 @@ void GatherQQMM::eval_gpu(const std::vector& inputs, array& out) { int K = x.shape(-1); int M = non_batched ? x.size() / K : x.shape(-2); int N = out.shape(-1); + + // temporary, until we add proper scaling for gather qqmm + if (global_scale_w) { + int E = w_q.size() / w_q.shape(-1) / w_q.shape(-2); + array gs_e(Shape{E}, float32, nullptr, {}); + broadcast(*global_scale_w, gs_e); + global_scale_w = ensure_row_contiguous(gs_e, d, s); + } + gather_qmv( x, w_q, diff --git a/mlx/ops.cpp b/mlx/ops.cpp index 34bf40538f..fb7979c34a 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -5577,9 +5577,14 @@ array gather_qmm( std::optional group_size_ /* = std::nullopt */, std::optional bits_ /* = std::nullopt */, const std::string& mode /* = "affine" */, + const std::optional& global_scale /* = std::nullopt */, bool sorted_indices /* = false */, StreamOrDevice s /* = {} */) { if (!lhs_indices_ && !rhs_indices_) { + if (global_scale) { + throw std::invalid_argument( + "[gather_qmm] Global scale is not supported without indices."); + } return quantized_matmul( x, w, scales, biases, transpose, group_size_, bits_, mode, s); } @@ -5590,6 +5595,32 @@ array gather_qmm( quantization_params_from_mode(qmode, group_size_, bits_); auto [w_inner_dims, w_outer_dims] = extract_quantized_matmul_dims( "gather_qmm", x, w, scales, biases, transpose, group_size, bits); + if (global_scale) { + if (qmode != QuantizationMode::Nvfp4) { + throw std::invalid_argument( + "[gather_qmm] Global scale is only supported for 'nvfp4' " + "quantization mode."); + } + if (global_scale->dtype() != float32) { + std::ostringstream msg; + msg << "[gather_qmm] Global scale must have dtype float32 but got " + << global_scale->dtype() << "."; + throw std::invalid_argument(msg.str()); + } + // One scale per expert, so it matches the batch dimensions of w. + Shape expected(w.shape().begin(), w.shape().end() - 2); + if (global_scale->shape() != expected) { + std::ostringstream msg; + msg << "[gather_qmm] Global scale must have one entry per expert with " + << "shape " << expected << " but got " << global_scale->shape() + << "."; + throw std::invalid_argument(msg.str()); + } + if (!mx::metal::is_available) { + throw std::invalid_argument( + "[gather_qmm] Global scale is not supported on Metal Backend."); + } + } if (qmode == QuantizationMode::Affine) { out_type = promote_types(x.dtype(), out_type); } else { @@ -5623,12 +5654,12 @@ array gather_qmm( std::move(lhs_indices), std::move(rhs_indices)}; } else { - inputs = { - astype(x, out_type, s), - std::move(w), - std::move(scales), - std::move(lhs_indices), - std::move(rhs_indices)}; + inputs = {astype(x, out_type, s), std::move(w), std::move(scales)}; + if (global_scale) { + inputs.push_back(*global_scale); + } + inputs.push_back(std::move(lhs_indices)); + inputs.push_back(std::move(rhs_indices)); } return array( std::move(out_shape), diff --git a/mlx/ops.h b/mlx/ops.h index f597753b1e..63c84e6cf4 100644 --- a/mlx/ops.h +++ b/mlx/ops.h @@ -1626,6 +1626,7 @@ MLX_API array gather_qmm( std::optional group_size = std::nullopt, std::optional bits = std::nullopt, const std::string& mode = "affine", + const std::optional& global_scale = std::nullopt, bool sorted_indices = false, StreamOrDevice s = {}); diff --git a/mlx/primitives.cpp b/mlx/primitives.cpp index 7a6c729c39..922d22bfbb 100644 --- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -3722,6 +3722,9 @@ std::vector GatherQMM::vjp( auto biases = (mode_ == QuantizationMode::Affine) ? std::optional(primals[3]) : std::nullopt; + auto global_scale = (mode_ != QuantizationMode::Affine && primals.size() == 6) + ? std::optional(primals[3]) + : std::nullopt; int M = cotan.shape(-2); int K = x.shape(-1); @@ -3744,6 +3747,7 @@ std::vector GatherQMM::vjp( group_size_, bits_, quantization_mode_to_string(mode_), + global_scale, sorted, stream()); if (sorted && no_broadcast) { diff --git a/python/src/ops.cpp b/python/src/ops.cpp index 677d998af1..87551c6b90 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -4823,11 +4823,12 @@ void init_ops(nb::module_& m) { "group_size"_a = nb::none(), "bits"_a = nb::none(), "mode"_a = "affine", + "global_scale"_a = nb::none(), nb::kw_only(), "sorted_indices"_a = false, "stream"_a = nb::none(), nb::sig( - "def gather_qmm(x: array, w: array, /, scales: array, biases: array | None = None, lhs_indices: array | None = None, rhs_indices: array | None = None, transpose: bool = True, group_size: int | None = None, bits: int | None = None, mode: str = 'affine', *, sorted_indices: bool = False, stream: StreamOrDevice = None) -> array"), + "def gather_qmm(x: array, w: array, /, scales: array, biases: array | None = None, lhs_indices: array | None = None, rhs_indices: array | None = None, transpose: bool = True, group_size: int | None = None, bits: int | None = None, mode: str = 'affine', global_scale: array | None = None, *, sorted_indices: bool = False, stream: StreamOrDevice = None) -> array"), R"pbdoc( Perform quantized matrix multiplication with matrix-level gather. @@ -4857,6 +4858,9 @@ void init_ops(nb::module_& m) { ``w`` in the quantized array. See supported values and defaults in the :ref:`table of quantization modes `. Default: ``None``. mode (str, optional): The quantization mode. Default: ``"affine"``. + global_scale (array, optional): The per-input float32 scale used for + ``nvfp4`` quantization of ``w``. Only supported on Metal. + Default: ``None``. sorted_indices (bool, optional): May allow a faster implementation if the passed indices are sorted. Default: ``False``. diff --git a/python/tests/test_quantized.py b/python/tests/test_quantized.py index 319e9d2d67..bb69afd517 100644 --- a/python/tests/test_quantized.py +++ b/python/tests/test_quantized.py @@ -1693,6 +1693,83 @@ def gmm(s, x, wq): ds = mx.grad(gmm)(s, x, wq) + @unittest.skipIf( + not mx.metal.is_available(), "Global scale is only supported on Metal backend" + ) + def test_gather_qmm_global_scale(self): + mx.random.seed(0) + N, K = 128, 256 + + def rel_err(got, expected): + # Per row + d = mx.abs(got.astype(mx.float32) - expected.astype(mx.float32)) + scale = mx.abs(expected.astype(mx.float32)) + axes = tuple(range(1, d.ndim)) + return (d.max(axis=axes) / mx.maximum(scale.max(axis=axes), 1e-20)).max() + + def quantize_experts(w): + """One tensor scale per expert.""" + E = w.shape[0] + gs = mx.stack([mx.abs(w[e]).max().astype(mx.float32) for e in range(E)]) + qs = [mx.quantize(w[e], mode="nvfp4", global_scale=gs[e]) for e in range(E)] + w_hat = mx.stack( + [ + mx.dequantize( + q, sc, mode="nvfp4", global_scale=gs[e], dtype=w.dtype + ) + for e, (q, sc) in enumerate(qs) + ] + ) + return ( + mx.stack([q for q, _ in qs]), + mx.stack([sc for _, sc in qs]), + gs, + w_hat, + ) + + for E in [4, 128, 234]: + # Each expert has a different scale, so we multiply by a factor + # to make the experts have different magnitudes. + factors = mx.array( + [[1, 2, 4, 8][e % 4] for e in range(E)], mx.float32 + ).reshape((E, 1, 1)) + + for dtype in [mx.float32, mx.float16, mx.bfloat16]: + for transpose in [True, False]: + wshape = (E, N, K) if transpose else (E, K, N) + w = (mx.random.normal(wshape) * factors).astype(dtype) + wq, s, gs, w_hat = quantize_experts(w) + self.assertEqual(gs.shape, (E,)) + + for M, B, sort in [(32, 2, False), (1, 2, False), (256, 4, True)]: + with self.subTest( + E=E, dtype=dtype, transpose=transpose, M=M, B=B + ): + x = mx.random.normal((B, M, K)).astype(dtype) + indices = mx.random.randint(0, E, (B,)) + if sort: + indices = mx.sort(indices) + + wg = w_hat[indices] + expected = x @ (wg.swapaxes(-1, -2) if transpose else wg) + kwargs = dict( + rhs_indices=indices, + transpose=transpose, + mode="nvfp4", + sorted_indices=sort, + ) + + out = mx.gather_qmm(x, wq, s, global_scale=gs, **kwargs) + tol = 1e-5 if dtype == mx.float32 else 3e-2 + self.assertLess(rel_err(out, expected), tol) + + # Each expert uses its own scale, not a neighbour's + rotated = mx.concatenate([gs[1:], gs[:1]]) + wrong = mx.gather_qmm( + x, wq, s, global_scale=rotated, **kwargs + ) + self.assertGreater(rel_err(wrong, expected), 0.5) + def test_quantize_strided(self): N = 64 mode = "nvfp4" From f7c34049e7514928844d1ac0464bcbb0033a840f Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Fri, 4 Sep 2026 01:45:14 +0200 Subject: [PATCH 2/4] fix metal::is_available --- mlx/ops.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index fb7979c34a..05393d8ea9 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -5616,7 +5616,7 @@ array gather_qmm( << "."; throw std::invalid_argument(msg.str()); } - if (!mx::metal::is_available) { + if (!metal::is_available) { throw std::invalid_argument( "[gather_qmm] Global scale is not supported on Metal Backend."); } From 285ab899087ecfbae9390c1e96b1571f86a5ef8a Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Fri, 4 Sep 2026 02:13:48 +0200 Subject: [PATCH 3/4] fix shapes and errors --- mlx/ops.cpp | 4 ++-- mlx/primitives.cpp | 11 +++++++---- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/mlx/ops.cpp b/mlx/ops.cpp index c68308f454..1a41688d28 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -5575,9 +5575,9 @@ array gather_qmm( << "."; throw std::invalid_argument(msg.str()); } - if (!metal::is_available) { + if (to_stream(s).device != Device::gpu || !metal::is_available()) { throw std::invalid_argument( - "[gather_qmm] Global scale is not supported on Metal Backend."); + "[gather_qmm] Global scale is only supported on the Metal backend."); } } if (qmode == QuantizationMode::Affine) { diff --git a/mlx/primitives.cpp b/mlx/primitives.cpp index 0de0ea1ae0..9f810230ad 100644 --- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -3738,6 +3738,7 @@ std::vector GatherQMM::vjp( int M = cotan.shape(-2); int K = x.shape(-1); + int first_index_arg = primals.size() - 2; bool sorted = left_sorted_ || right_sorted_; bool no_broadcast = rhs_indices.size() * M * K == x.size(); @@ -3776,15 +3777,18 @@ std::vector GatherQMM::vjp( } // gradient wrt to the indices is undefined - else if (arg > 3) { + else if (arg >= first_index_arg) { throw std::runtime_error( "[GatherQMM::vjp] cannot compute the gradient wrt the indices."); } - // gradient wrt to w_q, scales or biases + // gradient wrt to w_q, scales, biases or the global scale else if (arg == 1) { throw std::runtime_error( "[GatherQMM::vjp] no gradient wrt the quantized weights."); + } else if (global_scale && arg == 3) { + throw std::runtime_error( + "[GatherQMM::vjp] no gradient wrt the global scale."); } else { if (mode_ != QuantizationMode::Affine) { std::ostringstream msg; @@ -3856,8 +3860,7 @@ bool GatherQMM::is_equivalent(const Primitive& other) const { std::vector GatherQMM::output_shapes(const std::vector& inputs) { const auto& x = inputs[0]; const auto& w = inputs[1]; - const auto& lhs_indices = - (mode_ == QuantizationMode::Affine) ? inputs[4] : inputs[3]; + const auto& lhs_indices = inputs[inputs.size() - 2]; int w_outer = transpose_ ? w.shape(-2) : w.shape(-1) * 32 / bits_; auto out_shape = lhs_indices.shape(); out_shape.push_back(x.shape(-2)); From af85e41778845f0f7f4f340da0c6d57bfd1c69b1 Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Tue, 8 Sep 2026 15:17:14 +0200 Subject: [PATCH 4/4] address comments --- mlx/backend/metal/kernels/quantized.h | 1 + mlx/backend/metal/kernels/quantized_nax.h | 3 +- mlx/backend/metal/quantized.cpp | 94 ++++++++--------------- python/tests/test_quantized.py | 82 ++++++++++---------- 4 files changed, 77 insertions(+), 103 deletions(-) diff --git a/mlx/backend/metal/kernels/quantized.h b/mlx/backend/metal/kernels/quantized.h index 42855e14d8..15f044404f 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -2276,6 +2276,7 @@ template < const int group_size, const int bits, const bool aligned_N, + const bool has_global_scale = false, const int BM = 32, const int BK = 32, const int BN = 32> diff --git a/mlx/backend/metal/kernels/quantized_nax.h b/mlx/backend/metal/kernels/quantized_nax.h index 8e02a7e822..c89e66b7d9 100644 --- a/mlx/backend/metal/kernels/quantized_nax.h +++ b/mlx/backend/metal/kernels/quantized_nax.h @@ -1320,7 +1320,8 @@ template < const int BK = 64, const int BN = 64, const int WM = 2, - const int WN = 2> + const int WN = 2, + const bool has_global_scale = false> [[kernel]] void affine_gather_qmm_t_nax( const device uint32_t* w [[buffer(0)]], const device T* scales [[buffer(1)]], diff --git a/mlx/backend/metal/quantized.cpp b/mlx/backend/metal/quantized.cpp index 84e5e0d9db..c8f8d9c443 100644 --- a/mlx/backend/metal/quantized.cpp +++ b/mlx/backend/metal/quantized.cpp @@ -972,35 +972,21 @@ void gather_qmm_nax( global_scale ? "_hgs" : ""); MTL::ComputePipelineState* kernel; if (transpose) { - kernel = global_scale ? get_qmm_nax_kernel_wrapped( - d, - kname, - "gather_qmm_t_nax_", - mode, - type_string, - group_size, - bits, - aligned, - bm, - bk, - bn, - wm, - wn, - true) - : get_qmm_nax_kernel_wrapped( - d, - kname, - "gather_qmm_t_nax_", - mode, - type_string, - group_size, - bits, - aligned, - bm, - bk, - bn, - wm, - wn); + kernel = get_qmm_nax_kernel_wrapped( + d, + kname, + "gather_qmm_t_nax", + mode, + type_string, + group_size, + bits, + aligned, + bm, + bk, + bn, + wm, + wn, + global_scale.has_value()); } else { kernel = get_qmm_nax_kernel_wrapped( d, @@ -1302,38 +1288,26 @@ void gather_qmm( global_scale ? "_hgs" : ""); MTL::ComputePipelineState* kernel; if (transpose) { - kernel = global_scale ? get_quantized_kernel_wrapped( - d, - kname, - "gather_qmm_t", - mode, - type_string, - group_size, - bits, - aligned, - true) - : get_quantized_kernel_wrapped( - d, - kname, - "gather_qmm_t", - mode, - type_string, - group_size, - bits, - aligned); + kernel = get_quantized_kernel_wrapped( + d, + kname, + "gather_qmm_t", + mode, + type_string, + group_size, + bits, + aligned, + global_scale.has_value()); } else { - kernel = global_scale - ? get_quantized_kernel_wrapped( - d, - kname, - "gather_qmm_n", - mode, - type_string, - group_size, - bits, - true) - : get_quantized_kernel_wrapped( - d, kname, "gather_qmm_n", mode, type_string, group_size, bits); + kernel = get_quantized_kernel_wrapped( + d, + kname, + "gather_qmm_n", + mode, + type_string, + group_size, + bits, + global_scale.has_value()); } auto& compute_encoder = metal::get_command_encoder(s); diff --git a/python/tests/test_quantized.py b/python/tests/test_quantized.py index 6f89ba0710..17a2fa4a23 100644 --- a/python/tests/test_quantized.py +++ b/python/tests/test_quantized.py @@ -1848,48 +1848,46 @@ def quantize_experts(w): w_hat, ) - for E in [4, 128, 234]: - # Each expert has a different scale, so we multiply by a factor - # to make the experts have different magnitudes. - factors = mx.array( - [[1, 2, 4, 8][e % 4] for e in range(E)], mx.float32 - ).reshape((E, 1, 1)) - - for dtype in [mx.float32, mx.float16, mx.bfloat16]: - for transpose in [True, False]: - wshape = (E, N, K) if transpose else (E, K, N) - w = (mx.random.normal(wshape) * factors).astype(dtype) - wq, s, gs, w_hat = quantize_experts(w) - self.assertEqual(gs.shape, (E,)) - - for M, B, sort in [(32, 2, False), (1, 2, False), (256, 4, True)]: - with self.subTest( - E=E, dtype=dtype, transpose=transpose, M=M, B=B - ): - x = mx.random.normal((B, M, K)).astype(dtype) - indices = mx.random.randint(0, E, (B,)) - if sort: - indices = mx.sort(indices) - - wg = w_hat[indices] - expected = x @ (wg.swapaxes(-1, -2) if transpose else wg) - kwargs = dict( - rhs_indices=indices, - transpose=transpose, - mode="nvfp4", - sorted_indices=sort, - ) - - out = mx.gather_qmm(x, wq, s, global_scale=gs, **kwargs) - tol = 1e-5 if dtype == mx.float32 else 3e-2 - self.assertLess(rel_err(out, expected), tol) - - # Each expert uses its own scale, not a neighbour's - rotated = mx.concatenate([gs[1:], gs[:1]]) - wrong = mx.gather_qmm( - x, wq, s, global_scale=rotated, **kwargs - ) - self.assertGreater(rel_err(wrong, expected), 0.5) + tests = product( + [4, 128, 234], # E + [mx.float32, mx.float16, mx.bfloat16], # dtype + [True, False], # transpose + [(32, 2, False), (1, 2, False), (256, 4, True)], # M, B, sort + ) + for E, dtype, transpose, (M, B, sort) in tests: + with self.subTest(E=E, dtype=dtype, transpose=transpose, M=M, B=B): + # Each expert has a different scale, so we multiply by a factor + # to make the experts have different magnitudes. + factors = mx.array( + [[1, 2, 4, 8][e % 4] for e in range(E)], mx.float32 + ).reshape((E, 1, 1)) + wshape = (E, N, K) if transpose else (E, K, N) + w = (mx.random.normal(wshape) * factors).astype(dtype) + wq, s, gs, w_hat = quantize_experts(w) + self.assertEqual(gs.shape, (E,)) + + x = mx.random.normal((B, M, K)).astype(dtype) + indices = mx.random.randint(0, E, (B,)) + if sort: + indices = mx.sort(indices) + + wg = w_hat[indices] + expected = x @ (wg.swapaxes(-1, -2) if transpose else wg) + kwargs = dict( + rhs_indices=indices, + transpose=transpose, + mode="nvfp4", + sorted_indices=sort, + ) + + out = mx.gather_qmm(x, wq, s, global_scale=gs, **kwargs) + tol = 1e-5 if dtype == mx.float32 else 3e-2 + self.assertLess(rel_err(out, expected), tol) + + # Each expert uses its own scale, not a neighbour's + rotated = mx.concatenate([gs[1:], gs[:1]]) + wrong = mx.gather_qmm(x, wq, s, global_scale=rotated, **kwargs) + self.assertGreater(rel_err(wrong, expected), 0.5) def test_quantize_strided(self): N = 64