mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-17 20:31:47 +02:00
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) <noreply@anthropic.com>
This commit is contained in:
committed by
Ruben Ortlam
co-authored by
Claude Opus 5
parent
c9d34ea9fd
commit
ae6349d40e
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user