From 2232bc8b5f8cd91493a2c251b077c4d45cf2e832 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Thu, 1 Oct 2026 08:19:24 +0300 Subject: [PATCH] metal : use bf16 math for mxfp4 mul-mat (#29770) --- ggml/src/ggml-metal/kernels/dequantize.h | 12 ++++++++---- ggml/src/ggml-metal/kernels/mul_mm.metal | 14 ++++++++++---- tests/test-backend-ops.cpp | 3 ++- 3 files changed, 20 insertions(+), 9 deletions(-) diff --git a/ggml/src/ggml-metal/kernels/dequantize.h b/ggml/src/ggml-metal/kernels/dequantize.h index 0d1429d9d3..54e00faf6f 100644 --- a/ggml/src/ggml-metal/kernels/dequantize.h +++ b/ggml/src/ggml-metal/kernels/dequantize.h @@ -351,12 +351,16 @@ void dequantize_mxfp4(device const block_mxfp4 * xb, short il, thread type4x4 & const float d = e8m0_to_fp32(xb->e); const uint8_t shr = il >= 1 ? 4 : 0; + float4x4 reg_f; + for (int i = 0; i < 4; ++i) { - reg[i][0] = d * kvalues_mxfp4_f[(q2[4*i + 0] >> shr) & 0x0F]; - reg[i][1] = d * kvalues_mxfp4_f[(q2[4*i + 1] >> shr) & 0x0F]; - reg[i][2] = d * kvalues_mxfp4_f[(q2[4*i + 2] >> shr) & 0x0F]; - reg[i][3] = d * kvalues_mxfp4_f[(q2[4*i + 3] >> shr) & 0x0F]; + reg_f[i][0] = d * kvalues_mxfp4_f[(q2[4*i + 0] >> shr) & 0x0F]; + reg_f[i][1] = d * kvalues_mxfp4_f[(q2[4*i + 1] >> shr) & 0x0F]; + reg_f[i][2] = d * kvalues_mxfp4_f[(q2[4*i + 2] >> shr) & 0x0F]; + reg_f[i][3] = d * kvalues_mxfp4_f[(q2[4*i + 3] >> shr) & 0x0F]; } + + reg = (type4x4) reg_f; } template diff --git a/ggml/src/ggml-metal/kernels/mul_mm.metal b/ggml/src/ggml-metal/kernels/mul_mm.metal index a25838f926..9d262da1d5 100644 --- a/ggml/src/ggml-metal/kernels/mul_mm.metal +++ b/ggml/src/ggml-metal/kernels/mul_mm.metal @@ -226,7 +226,7 @@ kernel void kernel_mul_mm( const short ib = 8*sx + sy; - *(sa + 64*ib + 8*ly + lx) = loop_k + 16*il + i < args.ne00 ? *((device T0 *) x + i) : 0; + *(sa + 64*ib + 8*ly + lx) = loop_k + 16*il + i < args.ne00 ? (S0) *((device T0 *) x + i) : 0; } } else { S0_4x4 temp_a; @@ -582,8 +582,6 @@ kernel void kernel_mul_mm_id( const bool has_hi = nr1 > NR1H; - const short lb1 = (short) tiitg/NL1; // 0 .. NR1-1, this thread's row of the B tile - // power-of-two rescaling const float s1_inv = FC_mul_mm_id_amax ? ((device const float *) amax)[0] : 1.0f; const float s1_scale = FC_mul_mm_id_amax ? ((device const float *) amax)[1] : 1.0f; @@ -701,7 +699,7 @@ kernel void kernel_mul_mm_id( //const short lx = (tiitg/NL0)%8; //const short ly = i%8; - *(sa + NK*(8*sy + ly) + 8*sx + lx) = loop_k + 16*il + i < args.ne00 ? *((device T0 *) x + i) : 0; + *(sa + NK*(8*sy + ly) + 8*sx + lx) = loop_k + 16*il + i < args.ne00 ? (S0) *((device T0 *) x + i) : 0; } } else { S0_4x4 temp_a; @@ -923,7 +921,11 @@ template [[host_name("kernel_mul_mm_id_q4_1_f32")]] kernel mul_mm_id kernel_m template [[host_name("kernel_mul_mm_id_q5_0_f32")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q5_1_f32")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q8_0_f32")]] kernel mul_mm_id kernel_mul_mm_id; +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_mul_mm_id_mxfp4_f32")]] kernel mul_mm_id kernel_mul_mm_id; +#else template [[host_name("kernel_mul_mm_id_mxfp4_f32")]] kernel mul_mm_id kernel_mul_mm_id; +#endif template [[host_name("kernel_mul_mm_id_q2_K_f32")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q3_K_f32")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q4_K_f32")]] kernel mul_mm_id kernel_mul_mm_id; @@ -949,7 +951,11 @@ template [[host_name("kernel_mul_mm_id_q4_1_f16")]] kernel mul_mm_id kernel_m template [[host_name("kernel_mul_mm_id_q5_0_f16")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q5_1_f16")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q8_0_f16")]] kernel mul_mm_id kernel_mul_mm_id; +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_mul_mm_id_mxfp4_f16")]] kernel mul_mm_id kernel_mul_mm_id; +#else template [[host_name("kernel_mul_mm_id_mxfp4_f16")]] kernel mul_mm_id kernel_mul_mm_id; +#endif template [[host_name("kernel_mul_mm_id_q2_K_f16")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q3_K_f16")]] kernel mul_mm_id kernel_mul_mm_id; template [[host_name("kernel_mul_mm_id_q4_K_f16")]] kernel mul_mm_id kernel_mul_mm_id; diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5dc9cef3de..c38b00b5bd 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -5215,7 +5215,8 @@ static void init_mul_mat_id_tensors(ggml_context * ctx, int n_mats, float amax = for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { if (t->type == GGML_TYPE_I32) { continue; - } else if (amax != 1.0f && t->type == GGML_TYPE_F32) { + } + if (amax != 1.0f && t->type == GGML_TYPE_F32) { init_tensor_uniform(t, -amax, amax); } else { init_tensor_uniform(t);