From ae6349d40e8ef0235fa40653fdcc45a1c425d585 Mon Sep 17 00:00:00 2001 From: "Piotr Wilkin (ilintar)" Date: Wed, 16 Sep 2026 11:23:37 +0200 Subject: [PATCH] vulkan: add IQ4_XS support to the coopmat1 integer matmul shader (#28440) Adds IQ4_XS to mul_mmq_cm1: dedicated block_a_load/block_a_to_shmem that expand both nibbles of each packed32 word through cm1_kvalues, LOAD_VEC_A 8 and an IQ4_XS-sized a_panel_bytes estimate for the L2-friendly scheduling. Assisted-by: OpenAI Codex Co-authored-by: Claude Opus 5 (1M context) --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 4 ++- .../vulkan-shaders/mul_mmq_cm1.comp | 10 ++++-- .../vulkan-shaders/mul_mmq_cm1_funcs.glsl | 36 +++++++++++++++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 2 +- 4 files changed, 47 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 382e9ac8fe..89cff61011 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1531,7 +1531,7 @@ static bool ggml_vk_matmul_cm1_int_shmem_support(const vk_device& device, const case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_1: case GGML_TYPE_Q4_K: case GGML_TYPE_Q5_K: has_dm = true; break; - case GGML_TYPE_IQ4_NL: case GGML_TYPE_MXFP4: + case GGML_TYPE_IQ4_NL: case GGML_TYPE_IQ4_XS: case GGML_TYPE_MXFP4: has_kvalues = true; break; case GGML_TYPE_Q3_K: case GGML_TYPE_Q6_K: kscales2 = true; break; @@ -2458,6 +2458,7 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q5_1, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int, "matmul_q5_1_q8_1", matmul_q5_1_q8_1_cm1_len, matmul_q5_1_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); } cm1_create_mmq({GGML_TYPE_Q8_0, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int, "matmul_q8_0_q8_1", matmul_q8_0_q8_1_cm1_len, matmul_q8_0_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); cm1_create_mmq({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int, "matmul_iq4_nl_q8_1", matmul_iq4_nl_q8_1_cm1_len, matmul_iq4_nl_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); + cm1_create_mmq({GGML_TYPE_IQ4_XS, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int, "matmul_iq4_xs_q8_1", matmul_iq4_xs_q8_1_cm1_len, matmul_iq4_xs_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); cm1_create_mmq({GGML_TYPE_MXFP4, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int, "matmul_mxfp4_q8_1", matmul_mxfp4_q8_1_cm1_len, matmul_mxfp4_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); cm1_create_mmq({GGML_TYPE_Q3_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_q3_k_q8_1", matmul_q3_k_q8_1_cm1_len, matmul_q3_k_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q4_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int, "matmul_q4_k_q8_1", matmul_q4_k_q8_1_cm1_len, matmul_q4_k_q8_1_cm1_data, sizeof(vk_mat_mat_push_constants), 3); } @@ -2530,6 +2531,7 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { cm1_create_mmq({GGML_TYPE_Q5_1, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int, "matmul_id_subgroup_q5_1_q8_1", matmul_id_subgroup_q5_1_q8_1_cm1_len, matmul_id_subgroup_q5_1_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); cm1_create_mmq({GGML_TYPE_Q8_0, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int, "matmul_id_subgroup_q8_0_q8_1", matmul_id_subgroup_q8_0_q8_1_cm1_len, matmul_id_subgroup_q8_0_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); cm1_create_mmq({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int, "matmul_id_subgroup_iq4_nl_q8_1", matmul_id_subgroup_iq4_nl_q8_1_cm1_len, matmul_id_subgroup_iq4_nl_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); + cm1_create_mmq({GGML_TYPE_IQ4_XS, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int, "matmul_id_subgroup_iq4_xs_q8_1", matmul_id_subgroup_iq4_xs_q8_1_cm1_len, matmul_id_subgroup_iq4_xs_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); cm1_create_mmq({GGML_TYPE_MXFP4, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int, "matmul_id_subgroup_mxfp4_q8_1", matmul_id_subgroup_mxfp4_q8_1_cm1_len, matmul_id_subgroup_mxfp4_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); cm1_create_mmq({GGML_TYPE_Q3_K, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int_k, "matmul_id_subgroup_q3_k_q8_1", matmul_id_subgroup_q3_k_q8_1_cm1_len, matmul_id_subgroup_q3_k_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); cm1_create_mmq({GGML_TYPE_Q4_K, GGML_TYPE_Q8_1, true, false}, tc_mmq_cm1_int, "matmul_id_subgroup_q4_k_q8_1", matmul_id_subgroup_q4_k_q8_1_cm1_len, matmul_id_subgroup_q4_k_q8_1_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); 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 eccd69869b..7cab9a1195 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -110,11 +110,11 @@ shared float buf_b_d[BN * BK_STEP]; shared float buf_b_s[BN * BK_STEP]; #endif -#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS) || defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) shared int8_t cm1_kvalues[16]; #endif -#if defined(DATA_A_QUANT_K) || defined(DATA_A_NVFP4) +#if defined(DATA_A_QUANT_K) || defined(DATA_A_IQ4_XS) || defined(DATA_A_NVFP4) #define LOAD_VEC_A 8 #else #define LOAD_VEC_A (4 * QUANT_R) @@ -149,7 +149,7 @@ ACC_TYPE cm1_accumulate(ACC_TYPE prev, int acc_e, float scale_a, float nbias_a, #include "mul_mmq_cm1_funcs.glsl" void main() { -#if defined(DATA_A_IQ4_NL) +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS) if (gl_LocalInvocationIndex < 16u) { cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex]; } @@ -184,7 +184,11 @@ void main() { #else // L2-friendly workgroup scheduling const uint blocks_n = (p.N + BN - 1) / BN; +#if defined(DATA_A_IQ4_XS) + const uint a_panel_bytes = (BM * p.K) / 2 + (BM * p.K) / 32; +#else const uint a_panel_bytes = BM * p.K + (BM * p.K) / 16; +#endif const uint group_m = clamp(GROUP_A_BUDGET / max(a_panel_bytes, 1u), 1u, min(blocks_m, 32u)); const uint tiles_per_group = group_m * blocks_n; const uint lin = gl_WorkGroupID.y * blocks_m + (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 3b07934f98..1760c138e3 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 @@ -173,6 +173,42 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { } } +#elif defined(DATA_A_IQ4_XS) + +struct block_a_prefetch { + uint32_t qs; + float d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint ib32 = ib % 8; + blk.qs = data_a_packed32[ib_k].qs[4 * ib32 + loadr]; + blk.d = 0.0; + if (loadr == 0) { + const uint sl = (data_a_packed32[ib_k].scales_l >> (4 * ib32)) & 0xF; + const uint sh = (data_a_packed32[ib_k].scales_h >> (2 * ib32)) & 3; + blk.d = float(data_a_packed32[ib_k].d) * float(int(sl | (sh << 4)) - 32); + } + return blk; +} + +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(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(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] = blk.d; + } +} + #elif defined(DATA_A_MXFP4) struct block_a_prefetch { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 189b6a97b8..2a3f521c9c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -630,7 +630,7 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c } #endif - if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "mxfp4" + if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "iq4_xs" || tname == "mxfp4" || tname == "q3_k" || tname == "q4_k" || tname == "q5_k" || tname == "q6_k" || tname == "nvfp4")) { string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); }