Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion ggml/src/ggml-vulkan/ggml-vulkan.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4288,7 +4288,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;
Expand Down Expand Up @@ -5252,6 +5252,7 @@ static 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); }
Expand Down Expand Up @@ -5324,6 +5325,7 @@ static 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);
Expand Down
10 changes: 7 additions & 3 deletions ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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];
}
Expand Down Expand Up @@ -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);
Expand Down
36 changes: 36 additions & 0 deletions ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
2 changes: 1 addition & 1 deletion ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
Expand Down
Loading