support iq4_nl and mxfp4

This commit is contained in:
Ruben Ortlam
2026-08-29 11:51:15 +02:00
parent 3cd76432d7
commit dc08fa9101
3 changed files with 77 additions and 11 deletions
+14 -10
View File
@@ -4874,17 +4874,21 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
}
if (device->coopmat_int_support) {
CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, );
CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
}
GGML_ASSERT(device->subgroup_ballot);
@@ -153,6 +153,68 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
}
}
#elif defined(DATA_A_IQ4_NL)
struct block_a_prefetch {
uint32_t qs;
float16_t d;
};
block_a_prefetch block_a_load(uint ib, uint loadr) {
block_a_prefetch blk;
blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2],
data_a_packed16[ib].qs[loadr * 2 + 1]));
blk.d = data_a_packed16[ib].d;
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(kvalues_iq4nl_const[lo_idx.x], kvalues_iq4nl_const[lo_idx.y],
kvalues_iq4nl_const[lo_idx.z], kvalues_iq4nl_const[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]));
if (loadr == 0) {
buf_a_d[ks * BM + buf_ib] = float(blk.d);
}
}
#elif defined(DATA_A_MXFP4)
struct block_a_prefetch {
uint32_t qs;
uint8_t e;
};
block_a_prefetch block_a_load(uint ib, uint loadr) {
block_a_prefetch blk;
blk.qs = pack32(u8vec4(data_a[ib].qs[loadr * 4],
data_a[ib].qs[loadr * 4 + 1],
data_a[ib].qs[loadr * 4 + 2],
data_a[ib].qs[loadr * 4 + 3]));
blk.e = data_a[ib].e;
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(kvalues_mxfp4_const[lo_idx.x], kvalues_mxfp4_const[lo_idx.y],
kvalues_mxfp4_const[lo_idx.z], kvalues_mxfp4_const[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]));
if (loadr == 0) {
buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5;
}
}
#endif
// ===== B-side: load and store =====
@@ -629,7 +629,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")) {
if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "mxfp4")) {
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);
}
}