From ad46e9d244764f7c2ea12dc0bc91b2bb497200c2 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Mon, 24 Aug 2026 17:37:24 +0200 Subject: [PATCH] use shmem arrays for LUTs --- .../ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp | 16 ++++++++++++++++ .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 16 ++++++++-------- 2 files changed, 24 insertions(+), 8 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp index 7e79b0d80f..e9ffd608e5 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -101,6 +101,10 @@ shared float buf_b_s[BN * BK_STEP]; #define QUANT_OFFSET 16.0 #endif +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) +shared int8_t cm1_kvalues[16]; +#endif + #define LOAD_VEC_A (4 * QUANT_R) #define LOAD_VEC_B 16 @@ -119,6 +123,18 @@ shared int32_t cm_layout_probe[TM * TN]; #include "mul_mmq_cm1_funcs.glsl" void main() { +#if defined(DATA_A_IQ4_NL) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex]; + } + barrier(); +#elif defined(DATA_A_MXFP4) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex]; + } + barrier(); +#endif + const uint blocks_m = (p.M + BM - 1) / BM; const uint ik = gl_WorkGroupID.x / blocks_m; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl index 3fe104cf98..294dd75a2a 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl @@ -172,11 +172,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = - pack32(i8vec4(kvalues_iq4nl_const[lo_idx.x], kvalues_iq4nl_const[lo_idx.y], - kvalues_iq4nl_const[lo_idx.z], kvalues_iq4nl_const[lo_idx.w])); + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = - pack32(i8vec4(kvalues_iq4nl_const[hi_idx.x], kvalues_iq4nl_const[hi_idx.y], - kvalues_iq4nl_const[hi_idx.z], kvalues_iq4nl_const[hi_idx.w])); + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); if (loadr == 0) { buf_a_d[ks * BM + buf_ib] = float(blk.d); @@ -204,11 +204,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = - pack32(i8vec4(kvalues_mxfp4_const[lo_idx.x], kvalues_mxfp4_const[lo_idx.y], - kvalues_mxfp4_const[lo_idx.z], kvalues_mxfp4_const[lo_idx.w])); + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = - pack32(i8vec4(kvalues_mxfp4_const[hi_idx.x], kvalues_mxfp4_const[hi_idx.y], - kvalues_mxfp4_const[hi_idx.z], kvalues_mxfp4_const[hi_idx.w])); + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); if (loadr == 0) { buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5;