From f8a74a4aef6cb5cc553d250fd55feff51164bc00 Mon Sep 17 00:00:00 2001 From: Ruben Ortlam Date: Wed, 26 Aug 2026 04:52:28 +0200 Subject: [PATCH] remove LUT quants from unified shader --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 222 ++++--- .../vulkan-shaders/dequant_funcs_cm2.glsl | 22 +- .../vulkan-shaders/iq_shmem_init.glsl | 64 --- .../ggml-vulkan/vulkan-shaders/mul_mm.comp | 24 - .../vulkan-shaders/mul_mm_cm2.comp | 42 +- .../vulkan-shaders/mul_mm_funcs.glsl | 544 +++++++++--------- .../src/ggml-vulkan/vulkan-shaders/types.glsl | 51 +- .../vulkan-shaders/vulkan-shaders-gen.cpp | 32 +- 8 files changed, 485 insertions(+), 516 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index b7d4d09e5f..6a8689c415 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4642,14 +4642,27 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { return spec; }; - static const ggml_type quant_types[] = { + static const ggml_type non_lut_quant_types[] = { GGML_TYPE_Q1_0, GGML_TYPE_Q2_0, GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_Q8_0, GGML_TYPE_Q2_K, GGML_TYPE_TQ2_0, GGML_TYPE_Q3_K, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K, - GGML_TYPE_IQ1_S, GGML_TYPE_IQ1_M, GGML_TYPE_IQ2_XXS, GGML_TYPE_IQ2_XS, GGML_TYPE_IQ2_S, - GGML_TYPE_IQ3_XXS, GGML_TYPE_IQ3_S, GGML_TYPE_IQ4_XS, GGML_TYPE_IQ4_NL, - GGML_TYPE_MXFP4, GGML_TYPE_NVFP4, }; +#define FOR_EACH_LUT_TYPE(X) \ + X(GGML_TYPE_IQ1_S, iq1_s) \ + X(GGML_TYPE_IQ1_M, iq1_m) \ + X(GGML_TYPE_IQ2_XXS, iq2_xxs) \ + X(GGML_TYPE_IQ2_XS, iq2_xs) \ + X(GGML_TYPE_IQ2_S, iq2_s) \ + X(GGML_TYPE_IQ3_XXS, iq3_xxs) \ + X(GGML_TYPE_IQ3_S, iq3_s) \ + X(GGML_TYPE_IQ4_XS, iq4_xs) \ + X(GGML_TYPE_IQ4_NL, iq4_nl) \ + X(GGML_TYPE_MXFP4, mxfp4) \ + X(GGML_TYPE_NVFP4, nvfp4) +#define FOR_EACH_LUT_FP4_TYPE(X) \ + X(GGML_TYPE_MXFP4, mxfp4) \ + X(GGML_TYPE_NVFP4, nvfp4) + const int mul_mat_id_param_count = 5; using spec_fn_t = std::function(const std::vector&, bool)>; @@ -4740,20 +4753,32 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { create_mm_pipelines({GGML_TYPE_BF16, GGML_TYPE_BF16, false, false}, tc_mm, "matmul_bf16", matmul_bf16_cm2_len, matmul_bf16_cm2_data, sizeof(vk_mat_mat_push_constants), 3, cm2_spec, true); } #endif - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { auto& tc = ((type >= GGML_TYPE_Q2_K && type <= GGML_TYPE_Q6_K) || type == GGML_TYPE_TQ2_0) ? tc_mmq_k : tc_mmq; spec_fn_t qs = [&, type](const std::vector& wt, bool a) { return ggml_vk_mul_mm_cm2_spec(wt, a, (uint32_t)type); }; -#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) - if (device->ocp_fp4 && (type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4)) { - create_mm_pipelines({type, GGML_TYPE_F16, false, true}, tc, "matmul_quant_f16_ocp_f16acc", matmul_quant_f16_ocp_f16acc_cm2_len, matmul_quant_f16_ocp_f16acc_cm2_data, sizeof(vk_mat_mat_push_constants), 3, qs, true); - create_mm_pipelines({type, GGML_TYPE_F16, false, false}, tc, "matmul_quant_f16_ocp", matmul_quant_f16_ocp_cm2_len, matmul_quant_f16_ocp_cm2_data, sizeof(vk_mat_mat_push_constants), 3, qs, true); - } else -#endif - { - create_mm_pipelines({type, GGML_TYPE_F16, false, true}, tc, "matmul_quant_f16_f16acc", matmul_quant_f16_f16acc_cm2_len, matmul_quant_f16_f16acc_cm2_data, sizeof(vk_mat_mat_push_constants), 3, qs, true); - create_mm_pipelines({type, GGML_TYPE_F16, false, false}, tc, "matmul_quant_f16", matmul_quant_f16_cm2_len, matmul_quant_f16_cm2_data, sizeof(vk_mat_mat_push_constants), 3, qs, true); - } + create_mm_pipelines({type, GGML_TYPE_F16, false, true}, tc, "matmul_quant_f16_f16acc", matmul_quant_f16_f16acc_cm2_len, matmul_quant_f16_f16acc_cm2_data, sizeof(vk_mat_mat_push_constants), 3, qs, true); + create_mm_pipelines({type, GGML_TYPE_F16, false, false}, tc, "matmul_quant_f16", matmul_quant_f16_cm2_len, matmul_quant_f16_cm2_data, sizeof(vk_mat_mat_push_constants), 3, qs, true); } +#define X_CM2(TYPE, tstr) \ + { auto tc = filter_tc(tc_mmq, TYPE, false); \ + if (!tc.empty()) { \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, false, true}, tc, "matmul_" #tstr "_f16_f16acc", matmul_##tstr##_f16_f16acc_cm2_len, matmul_##tstr##_f16_f16acc_cm2_data, sizeof(vk_mat_mat_push_constants), 3, cm2_spec, true); \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, false, false}, tc, "matmul_" #tstr "_f16", matmul_##tstr##_f16_cm2_len, matmul_##tstr##_f16_cm2_data, sizeof(vk_mat_mat_push_constants), 3, cm2_spec, true); \ + } } + FOR_EACH_LUT_TYPE(X_CM2) +#undef X_CM2 +#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) + if (device->ocp_fp4) { +#define X_CM2_OCP(TYPE, tstr) \ + { auto tc = filter_tc(tc_mmq, TYPE, false); \ + if (!tc.empty()) { \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, false, true}, tc, "matmul_" #tstr "_f16_ocp_f16acc", matmul_##tstr##_f16_ocp_f16acc_cm2_len, matmul_##tstr##_f16_ocp_f16acc_cm2_data, sizeof(vk_mat_mat_push_constants), 3, cm2_spec, true); \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, false, false}, tc, "matmul_" #tstr "_f16_ocp", matmul_##tstr##_f16_ocp_cm2_len, matmul_##tstr##_f16_ocp_cm2_data, sizeof(vk_mat_mat_push_constants), 3, cm2_spec, true); \ + } } + FOR_EACH_LUT_FP4_TYPE(X_CM2_OCP) +#undef X_CM2_OCP + } +#endif GGML_ASSERT(device->subgroup_ballot); @@ -4765,19 +4790,31 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { create_mm_pipelines({GGML_TYPE_BF16, GGML_TYPE_BF16, true, false}, tc_mm, "matmul_id_subgroup_bf16", matmul_id_subgroup_bf16_cm2_len, matmul_id_subgroup_bf16_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, cm2_spec, true); } #endif - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { spec_fn_t qs_id = [&, type](const std::vector& wt, bool a) { return ggml_vk_mul_mm_cm2_spec(wt, a, (uint32_t)type); }; -#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) - if (device->ocp_fp4 && (type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4)) { - create_mm_pipelines({type, GGML_TYPE_F16, true, true}, tc_mmqid, "matmul_id_subgroup_quant_f16_ocp_f16acc", matmul_id_subgroup_quant_f16_ocp_f16acc_cm2_len, matmul_id_subgroup_quant_f16_ocp_f16acc_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, qs_id, true); - create_mm_pipelines({type, GGML_TYPE_F16, true, false}, tc_mmqid, "matmul_id_subgroup_quant_f16_ocp", matmul_id_subgroup_quant_f16_ocp_cm2_len, matmul_id_subgroup_quant_f16_ocp_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, qs_id, true); - } else -#endif - { - create_mm_pipelines({type, GGML_TYPE_F16, true, true}, tc_mmqid, "matmul_id_subgroup_quant_f16_f16acc", matmul_id_subgroup_quant_f16_f16acc_cm2_len, matmul_id_subgroup_quant_f16_f16acc_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, qs_id, true); - create_mm_pipelines({type, GGML_TYPE_F16, true, false}, tc_mmqid, "matmul_id_subgroup_quant_f16", matmul_id_subgroup_quant_f16_cm2_len, matmul_id_subgroup_quant_f16_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, qs_id, true); - } + create_mm_pipelines({type, GGML_TYPE_F16, true, true}, tc_mmqid, "matmul_id_subgroup_quant_f16_f16acc", matmul_id_subgroup_quant_f16_f16acc_cm2_len, matmul_id_subgroup_quant_f16_f16acc_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, qs_id, true); + create_mm_pipelines({type, GGML_TYPE_F16, true, false}, tc_mmqid, "matmul_id_subgroup_quant_f16", matmul_id_subgroup_quant_f16_cm2_len, matmul_id_subgroup_quant_f16_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, qs_id, true); } +#define X_CM2_ID(TYPE, tstr) \ + { auto tc = filter_tc(tc_mmqid, TYPE, true); \ + if (!tc.empty()) { \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, true, true}, tc, "matmul_id_subgroup_" #tstr "_f16_f16acc", matmul_id_subgroup_##tstr##_f16_f16acc_cm2_len, matmul_id_subgroup_##tstr##_f16_f16acc_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, cm2_spec, true); \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, true, false}, tc, "matmul_id_subgroup_" #tstr "_f16", matmul_id_subgroup_##tstr##_f16_cm2_len, matmul_id_subgroup_##tstr##_f16_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, cm2_spec, true); \ + } } + FOR_EACH_LUT_TYPE(X_CM2_ID) +#undef X_CM2_ID +#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) + if (device->ocp_fp4) { +#define X_CM2_ID_OCP(TYPE, tstr) \ + { auto tc = filter_tc(tc_mmqid, TYPE, true); \ + if (!tc.empty()) { \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, true, true}, tc, "matmul_id_subgroup_" #tstr "_f16_ocp_f16acc", matmul_id_subgroup_##tstr##_f16_ocp_f16acc_cm2_len, matmul_id_subgroup_##tstr##_f16_ocp_f16acc_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, cm2_spec, true); \ + create_mm_pipelines({TYPE, GGML_TYPE_F16, true, false}, tc, "matmul_id_subgroup_" #tstr "_f16_ocp", matmul_id_subgroup_##tstr##_f16_ocp_cm2_len, matmul_id_subgroup_##tstr##_f16_ocp_cm2_data, sizeof(vk_mat_mat_id_push_constants), 5, cm2_spec, true); \ + } } + FOR_EACH_LUT_FP4_TYPE(X_CM2_ID_OCP) +#undef X_CM2_ID_OCP + } +#endif } else #endif // defined(VK_NV_cooperative_matrix2) && defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT) #if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) @@ -4811,26 +4848,36 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { cm1_create({GGML_TYPE_BF16, GGML_TYPE_BF16, false, false}, tc_mm, "matmul_bf16", matmul_bf16_cm1_len, matmul_bf16_cm1_data, sizeof(vk_mat_mat_push_constants), 3); } #endif - for (const auto type : quant_types) { -#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) - if (device->ocp_fp4 && (type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4)) { - if (device->coopmat_acc_f16_support) { - cm1_create_quant({type, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_quant_f32_ocp_f16acc", matmul_quant_f32_ocp_f16acc_cm1_len, matmul_quant_f32_ocp_f16acc_cm1_data, sizeof(vk_mat_mat_push_constants), 3); - } - if (device->coopmat_acc_f32_support) { - cm1_create_quant({type, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_quant_f32_ocp", matmul_quant_f32_ocp_cm1_len, matmul_quant_f32_ocp_cm1_data, sizeof(vk_mat_mat_push_constants), 3); - } - } else -#endif - { - if (device->coopmat_acc_f16_support) { - cm1_create_quant({type, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_quant_f32_f16acc", matmul_quant_f32_f16acc_cm1_len, matmul_quant_f32_f16acc_cm1_data, sizeof(vk_mat_mat_push_constants), 3); - } - if (device->coopmat_acc_f32_support) { - cm1_create_quant({type, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_quant_f32", matmul_quant_f32_cm1_len, matmul_quant_f32_cm1_data, sizeof(vk_mat_mat_push_constants), 3); - } + for (const auto type : non_lut_quant_types) { + if (device->coopmat_acc_f16_support) { + cm1_create_quant({type, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_quant_f32_f16acc", matmul_quant_f32_f16acc_cm1_len, matmul_quant_f32_f16acc_cm1_data, sizeof(vk_mat_mat_push_constants), 3); + } + if (device->coopmat_acc_f32_support) { + cm1_create_quant({type, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_quant_f32", matmul_quant_f32_cm1_len, matmul_quant_f32_cm1_data, sizeof(vk_mat_mat_push_constants), 3); } } +#define X_CM1(TYPE, tstr) \ + if (device->coopmat_acc_f16_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_" #tstr "_f32_f16acc", matmul_##tstr##_f32_f16acc_cm1_len, matmul_##tstr##_f32_f16acc_cm1_data, sizeof(vk_mat_mat_push_constants), 3); \ + } \ + if (device->coopmat_acc_f32_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_" #tstr "_f32", matmul_##tstr##_f32_cm1_len, matmul_##tstr##_f32_cm1_data, sizeof(vk_mat_mat_push_constants), 3); \ + } + FOR_EACH_LUT_TYPE(X_CM1) +#undef X_CM1 +#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) + if (device->ocp_fp4) { +#define X_CM1_OCP(TYPE, tstr) \ + if (device->coopmat_acc_f16_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_" #tstr "_f32_ocp_f16acc", matmul_##tstr##_f32_ocp_f16acc_cm1_len, matmul_##tstr##_f32_ocp_f16acc_cm1_data, sizeof(vk_mat_mat_push_constants), 3); \ + } \ + if (device->coopmat_acc_f32_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_" #tstr "_f32_ocp", matmul_##tstr##_f32_ocp_cm1_len, matmul_##tstr##_f32_ocp_cm1_data, sizeof(vk_mat_mat_push_constants), 3); \ + } + FOR_EACH_LUT_FP4_TYPE(X_CM1_OCP) +#undef X_CM1_OCP + } +#endif GGML_ASSERT(device->subgroup_ballot); @@ -4848,26 +4895,36 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { cm1_create({GGML_TYPE_BF16, GGML_TYPE_BF16, true, false}, tc_mm, "matmul_id_subgroup_bf16", matmul_id_subgroup_bf16_cm1_len, matmul_id_subgroup_bf16_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); } #endif - for (const auto type : quant_types) { -#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) - if (device->ocp_fp4 && (type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4)) { - if (device->coopmat_acc_f16_support) { - cm1_create_quant({type, GGML_TYPE_F32, true, true}, tc_mmq, "matmul_id_subgroup_quant_f32_ocp_f16acc", matmul_id_subgroup_quant_f32_ocp_f16acc_cm1_len, matmul_id_subgroup_quant_f32_ocp_f16acc_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); - } - if (device->coopmat_acc_f32_support) { - cm1_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmq, "matmul_id_subgroup_quant_f32_ocp", matmul_id_subgroup_quant_f32_ocp_cm1_len, matmul_id_subgroup_quant_f32_ocp_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); - } - } else -#endif - { - if (device->coopmat_acc_f16_support) { - cm1_create_quant({type, GGML_TYPE_F32, true, true}, tc_mmq, "matmul_id_subgroup_quant_f32_f16acc", matmul_id_subgroup_quant_f32_f16acc_cm1_len, matmul_id_subgroup_quant_f32_f16acc_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); - } - if (device->coopmat_acc_f32_support) { - cm1_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmq, "matmul_id_subgroup_quant_f32", matmul_id_subgroup_quant_f32_cm1_len, matmul_id_subgroup_quant_f32_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); - } + for (const auto type : non_lut_quant_types) { + if (device->coopmat_acc_f16_support) { + cm1_create_quant({type, GGML_TYPE_F32, true, true}, tc_mmq, "matmul_id_subgroup_quant_f32_f16acc", matmul_id_subgroup_quant_f32_f16acc_cm1_len, matmul_id_subgroup_quant_f32_f16acc_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); + } + if (device->coopmat_acc_f32_support) { + cm1_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmq, "matmul_id_subgroup_quant_f32", matmul_id_subgroup_quant_f32_cm1_len, matmul_id_subgroup_quant_f32_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); } } +#define X_CM1_ID(TYPE, tstr) \ + if (device->coopmat_acc_f16_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, true, true}, tc_mmq, "matmul_id_subgroup_" #tstr "_f32_f16acc", matmul_id_subgroup_##tstr##_f32_f16acc_cm1_len, matmul_id_subgroup_##tstr##_f32_f16acc_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); \ + } \ + if (device->coopmat_acc_f32_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, true, false}, tc_mmq, "matmul_id_subgroup_" #tstr "_f32", matmul_id_subgroup_##tstr##_f32_cm1_len, matmul_id_subgroup_##tstr##_f32_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); \ + } + FOR_EACH_LUT_TYPE(X_CM1_ID) +#undef X_CM1_ID +#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) + if (device->ocp_fp4) { +#define X_CM1_ID_OCP(TYPE, tstr) \ + if (device->coopmat_acc_f16_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, true, true}, tc_mmq, "matmul_id_subgroup_" #tstr "_f32_ocp_f16acc", matmul_id_subgroup_##tstr##_f32_ocp_f16acc_cm1_len, matmul_id_subgroup_##tstr##_f32_ocp_f16acc_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); \ + } \ + if (device->coopmat_acc_f32_support) { \ + cm1_create({TYPE, GGML_TYPE_F32, true, false}, tc_mmq, "matmul_id_subgroup_" #tstr "_f32_ocp", matmul_id_subgroup_##tstr##_f32_ocp_cm1_len, matmul_id_subgroup_##tstr##_f32_ocp_cm1_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); \ + } + FOR_EACH_LUT_FP4_TYPE(X_CM1_ID_OCP) +#undef X_CM1_ID_OCP + } +#endif } else #endif // defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) { @@ -4918,10 +4975,15 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { // BF16 - no dot2 sg_create({GGML_TYPE_BF16, GGML_TYPE_BF16, false, false}, tc_mm, "matmul_bf16", matmul_bf16_len, matmul_bf16_data, sizeof(vk_mat_mat_push_constants), 3); - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { sg_create_quant({type, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_quant_f32_f16acc", SPV_DOT2_F16ACC(matmul_quant_f32), sizeof(vk_mat_mat_push_constants), 3); sg_create_quant({type, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_quant_f32", SPV_DOT2(matmul_quant_f32), sizeof(vk_mat_mat_push_constants), 3); } + #define X_SG(TYPE, tstr) \ + sg_create({TYPE, GGML_TYPE_F32, false, true}, tc_mmq, "matmul_" #tstr "_f32_f16acc", SPV_DOT2_F16ACC(matmul_##tstr##_f32), sizeof(vk_mat_mat_push_constants), 3); \ + sg_create({TYPE, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_" #tstr "_f32", SPV_DOT2(matmul_##tstr##_f32), sizeof(vk_mat_mat_push_constants), 3); + FOR_EACH_LUT_TYPE(X_SG) +#undef X_SG #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) if (device->integer_dot_product) { @@ -4950,10 +5012,15 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { sg_create({GGML_TYPE_F16, GGML_TYPE_F32, true, false}, tc_id, "matmul_id_subgroup_f16_f32", SPV_DOT2(matmul_id_subgroup_f16_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size_16); // BF16 id - no dot2 sg_create({GGML_TYPE_BF16, GGML_TYPE_BF16, true, false}, tc_id, "matmul_id_subgroup_bf16", matmul_id_subgroup_bf16_len, matmul_id_subgroup_bf16_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size_16); - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { sg_create_quant({type, GGML_TYPE_F32, true, true}, tc_mmqid, "matmul_id_subgroup_quant_f32_f16acc", SPV_DOT2_F16ACC(matmul_id_subgroup_quant_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size); sg_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_subgroup_quant_f32", SPV_DOT2(matmul_id_subgroup_quant_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size); } + #define X_SG_ID_SUB(TYPE, tstr) \ + sg_create({TYPE, GGML_TYPE_F32, true, true}, tc_mmqid, "matmul_id_subgroup_" #tstr "_f32_f16acc", SPV_DOT2_F16ACC(matmul_id_subgroup_##tstr##_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size); \ + sg_create({TYPE, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_subgroup_" #tstr "_f32", SPV_DOT2(matmul_id_subgroup_##tstr##_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size); + FOR_EACH_LUT_TYPE(X_SG_ID_SUB) +#undef X_SG_ID_SUB #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) if (device->integer_dot_product) { std::vector tc_mmqid_int = {{s_warptile_mmqid_int, s_mmq_wg_denoms, s_align}, {m_warptile_mmqid_int, m_mmq_wg_denoms, m_align}, {l_warptile_mmqid_int, l_mmq_wg_denoms, l_align}}; @@ -4980,10 +5047,15 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { sg_create({GGML_TYPE_F16, GGML_TYPE_F32, true, false}, tc_mm, "matmul_id_f16_f32", SPV_DOT2(matmul_id_f16_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); // BF16 id - no dot2 sg_create({GGML_TYPE_BF16, GGML_TYPE_BF16, true, false}, tc_mm, "matmul_id_bf16", matmul_id_bf16_len, matmul_id_bf16_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { sg_create_quant({type, GGML_TYPE_F32, true, true}, tc_mmqid, "matmul_id_quant_f32_f16acc", SPV_DOT2_F16ACC(matmul_id_quant_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); sg_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_quant_f32", SPV_DOT2(matmul_id_quant_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); } + #define X_SG_ID(TYPE, tstr) \ + sg_create({TYPE, GGML_TYPE_F32, true, true}, tc_mmqid, "matmul_id_" #tstr "_f32_f16acc", SPV_DOT2_F16ACC(matmul_id_##tstr##_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); \ + sg_create({TYPE, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_" #tstr "_f32", SPV_DOT2(matmul_id_##tstr##_f32), sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); + FOR_EACH_LUT_TYPE(X_SG_ID) +#undef X_SG_ID #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) if (device->integer_dot_product) { std::vector tc_mmqid_int = {{s_warptile_mmqid_int, s_mmq_wg_denoms, s_align}, {m_warptile_mmqid_int, m_mmq_wg_denoms, m_align}, {l_warptile_mmqid_int, l_mmq_wg_denoms, l_align}}; @@ -5013,9 +5085,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { sg_create({GGML_TYPE_F16, GGML_TYPE_F32, false, false}, tc_mm, "matmul_f16_f32", matmul_f16_f32_fp32_len, matmul_f16_f32_fp32_data, sizeof(vk_mat_mat_push_constants), 3); sg_create({GGML_TYPE_BF16, GGML_TYPE_BF16, false, false}, tc_mm, "matmul_bf16", matmul_bf16_fp32_len, matmul_bf16_fp32_data, sizeof(vk_mat_mat_push_constants), 3); - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { sg_create_quant({type, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_quant_f32", matmul_quant_f32_fp32_len, matmul_quant_f32_fp32_data, sizeof(vk_mat_mat_push_constants), 3); } + #define X_SG_FP32(TYPE, tstr) \ + sg_create({TYPE, GGML_TYPE_F32, false, false}, tc_mmq, "matmul_" #tstr "_f32", matmul_##tstr##_f32_fp32_len, matmul_##tstr##_f32_fp32_data, sizeof(vk_mat_mat_push_constants), 3); + FOR_EACH_LUT_TYPE(X_SG_FP32) +#undef X_SG_FP32 #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) if (device->integer_dot_product) { @@ -5042,20 +5118,30 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { sg_create({GGML_TYPE_F16, GGML_TYPE_F16, true, false}, tc_id, "matmul_id_subgroup_f16", matmul_id_subgroup_f16_fp32_len, matmul_id_subgroup_f16_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size_16); sg_create({GGML_TYPE_F16, GGML_TYPE_F32, true, false}, tc_id, "matmul_id_subgroup_f16_f32", matmul_id_subgroup_f16_f32_fp32_len, matmul_id_subgroup_f16_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size_16); sg_create({GGML_TYPE_BF16, GGML_TYPE_BF16, true, false}, tc_id, "matmul_id_subgroup_bf16", matmul_id_subgroup_bf16_fp32_len, matmul_id_subgroup_bf16_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size_16); - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { sg_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_subgroup_quant_f32", matmul_id_subgroup_quant_f32_fp32_len, matmul_id_subgroup_quant_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size); } + #define X_SG_ID_SUB_FP32(TYPE, tstr) \ + sg_create({TYPE, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_subgroup_" #tstr "_f32", matmul_id_subgroup_##tstr##_f32_fp32_len, matmul_id_subgroup_##tstr##_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count, mul_mat_subgroup_size); + FOR_EACH_LUT_TYPE(X_SG_ID_SUB_FP32) +#undef X_SG_ID_SUB_FP32 } else { sg_create({GGML_TYPE_F32, GGML_TYPE_F32, true, false}, tc_mm, "matmul_id_f32_f32", matmul_id_f32_f32_fp32_len, matmul_id_f32_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); sg_create({GGML_TYPE_F16, GGML_TYPE_F16, true, false}, tc_mm, "matmul_id_f16", matmul_id_f16_fp32_len, matmul_id_f16_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); sg_create({GGML_TYPE_F16, GGML_TYPE_F32, true, false}, tc_mm, "matmul_id_f16_f32", matmul_id_f16_f32_fp32_len, matmul_id_f16_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); sg_create({GGML_TYPE_BF16, GGML_TYPE_BF16, true, false}, tc_mm, "matmul_id_bf16", matmul_id_bf16_fp32_len, matmul_id_bf16_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); - for (const auto type : quant_types) { + for (const auto type : non_lut_quant_types) { sg_create_quant({type, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_quant_f32", matmul_id_quant_f32_fp32_len, matmul_id_quant_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); } + #define X_SG_ID_FP32(TYPE, tstr) \ + sg_create({TYPE, GGML_TYPE_F32, true, false}, tc_mmqid, "matmul_id_" #tstr "_f32", matmul_id_##tstr##_f32_fp32_len, matmul_id_##tstr##_f32_fp32_data, sizeof(vk_mat_mat_id_push_constants), mul_mat_id_param_count); + FOR_EACH_LUT_TYPE(X_SG_ID_FP32) +#undef X_SG_ID_FP32 } } } +#undef FOR_EACH_LUT_TYPE +#undef FOR_EACH_LUT_FP4_TYPE // BF16 fallback for coopmat devices without bf16 coopmat support if ((device->coopmat2 || device->coopmat_support) #if defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl index 55d2c3912f..d32a30c86f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl @@ -804,7 +804,7 @@ f16vec4 dequantFuncQ6_K_v(const in decodeBufQ6_K bl, const in uint blockCoords[2 return f16vec4((vec4(qi) - vec4(32.0f)) * vec4(float(dscale))); } -#if defined(DATA_A_IQ1_S) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ1_S) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ1_S { block_iq1_s block; }; @@ -851,7 +851,7 @@ f16vec4 dequantFuncIQ1_S_v(const in decodeBufIQ1_S bl, const in uint blockCoords } #endif -#if defined(DATA_A_IQ1_M) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ1_M) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ1_M { block_iq1_m block; }; @@ -910,7 +910,7 @@ f16vec4 dequantFuncIQ1_M_v(const in decodeBufIQ1_M bl, const in uint blockCoords } #endif -#if defined(DATA_A_IQ2_XXS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ2_XXS) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ2_XXS { block_iq2_xxs block; }; @@ -972,7 +972,7 @@ f16vec4 dequantFuncIQ2_XXS_v(const in decodeBufIQ2_XXS bl, const in uint blockCo } #endif -#if defined(DATA_A_IQ2_XS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ2_XS) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ2_XS { block_iq2_xs block; }; @@ -1025,7 +1025,7 @@ f16vec4 dequantFuncIQ2_XS_v(const in decodeBufIQ2_XS bl, const in uint blockCoor } #endif -#if defined(DATA_A_IQ2_S) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ2_S) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ2_S { block_iq2_s block; }; @@ -1079,7 +1079,7 @@ f16vec4 dequantFuncIQ2_S_v(const in decodeBufIQ2_S bl, const in uint blockCoords } #endif -#if defined(DATA_A_IQ3_XXS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ3_XXS) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ3_XXS { block_iq3_xxs block; }; @@ -1138,7 +1138,7 @@ f16vec4 dequantFuncIQ3_XXS_v(const in decodeBufIQ3_XXS bl, const in uint blockCo } #endif -#if defined(DATA_A_IQ3_S) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ3_S) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ3_S { block_iq3_s block; }; @@ -1188,7 +1188,7 @@ f16vec4 dequantFuncIQ3_S_v(const in decodeBufIQ3_S bl, const in uint blockCoords } #endif -#if defined(DATA_A_IQ4_XS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ4_XS) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ4_XS { block_iq4_xs block; }; @@ -1238,7 +1238,7 @@ f16vec4 dequantFuncIQ4_XS_v(const in decodeBufIQ4_XS bl, const in uint blockCoor } #endif -#if defined(DATA_A_IQ4_NL) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ4_NL) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ4_NL { block_iq4_nl block; }; @@ -1279,7 +1279,7 @@ f16vec4 dequantFuncIQ4_NL_v(const in decodeBufIQ4_NL bl, const in uint blockCoor } #endif -#if defined(DATA_A_MXFP4) || defined(MULMAT_QUANT) +#if defined(DATA_A_MXFP4) layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufMXFP4 { block_mxfp4 block; }; @@ -1333,7 +1333,7 @@ f16vec4 dequantFuncMXFP4_v(const in decodeBufMXFP4 bl, const in uint blockCoords } #endif -#if defined(DATA_A_NVFP4) || defined(MULMAT_QUANT) +#if defined(DATA_A_NVFP4) layout(buffer_reference, std430, buffer_reference_align = 4) buffer decodeBufNVFP4 { block_nvfp4 block; }; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/iq_shmem_init.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/iq_shmem_init.glsl index 5fb2ba9ab8..12e50ee9eb 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/iq_shmem_init.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/iq_shmem_init.glsl @@ -1,66 +1,2 @@ void init_iq_shmem(uvec3 wgsize) { - if (MmTypeA == GGML_TYPE_IQ1_S || MmTypeA == GGML_TYPE_IQ1_M) { - [[unroll]] for (uint i = 0; i < iq1s_grid_const.length(); i += wgsize.x) { - uint idx = i + gl_LocalInvocationIndex.x; - if (iq1s_grid_const.length() % wgsize.x == 0 || idx < iq1s_grid_const.length()) { - u16vec2 g = unpack16(iq1s_grid_const[idx]); - iq1s_grid[2*idx+0] = g.x; - iq1s_grid[2*idx+1] = g.y; - } - } - barrier(); - } else if (MmTypeA == GGML_TYPE_IQ2_XXS) { - [[unroll]] for (uint i = 0; i < iq2_grid.length(); i += wgsize.x) { - if (iq2xxs_grid_const.length() % wgsize.x == 0 || i + gl_LocalInvocationIndex.x < iq2xxs_grid_const.length()) { - iq2_grid[i + gl_LocalInvocationIndex.x] = iq2xxs_grid_const[i + gl_LocalInvocationIndex.x]; - } - } - barrier(); - } else if (MmTypeA == GGML_TYPE_IQ2_XS) { - [[unroll]] for (uint i = 0; i < iq2_grid.length(); i += wgsize.x) { - if (iq2xs_grid_const.length() % wgsize.x == 0 || i + gl_LocalInvocationIndex.x < iq2xs_grid_const.length()) { - iq2_grid[i + gl_LocalInvocationIndex.x] = iq2xs_grid_const[i + gl_LocalInvocationIndex.x]; - } - } - barrier(); - } else if (MmTypeA == GGML_TYPE_IQ2_S) { - [[unroll]] for (uint i = 0; i < iq2_grid.length(); i += wgsize.x) { - if (iq2s_grid_const.length() % wgsize.x == 0 || i + gl_LocalInvocationIndex.x < iq2s_grid_const.length()) { - iq2_grid[i + gl_LocalInvocationIndex.x] = iq2s_grid_const[i + gl_LocalInvocationIndex.x]; - } - } - barrier(); - } else if (MmTypeA == GGML_TYPE_IQ3_XXS) { - [[unroll]] for (uint i = 0; i < iq3_grid.length(); i += wgsize.x) { - if (iq3xxs_grid_const.length() % wgsize.x == 0 || i + gl_LocalInvocationIndex.x < iq3xxs_grid_const.length()) { - iq3_grid[i + gl_LocalInvocationIndex.x] = iq3xxs_grid_const[i + gl_LocalInvocationIndex.x]; - } - } - barrier(); - } else if (MmTypeA == GGML_TYPE_IQ3_S) { - [[unroll]] for (uint i = 0; i < iq3_grid.length(); i += wgsize.x) { - if (iq3s_grid_const.length() % wgsize.x == 0 || i + gl_LocalInvocationIndex.x < iq3s_grid_const.length()) { - iq3_grid[i + gl_LocalInvocationIndex.x] = iq3s_grid_const[i + gl_LocalInvocationIndex.x]; - } - } - barrier(); - } else if (MmTypeA == GGML_TYPE_IQ4_XS || MmTypeA == GGML_TYPE_IQ4_NL) { - for (uint i = gl_LocalInvocationIndex.x; i < kvalues_iq4nl.length(); i += wgsize.x) { - kvalues_iq4nl[i] = FLOAT_TYPE(kvalues_iq4nl_const[i]); - } - barrier(); - } -#if !defined(USE_OCP_FP4) - else if (MmTypeA == GGML_TYPE_MXFP4 || MmTypeA == GGML_TYPE_NVFP4) { - for (uint i = gl_LocalInvocationIndex.x; i < kvalues_mxfp4.length(); i += wgsize.x) { - kvalues_mxfp4[i] = kvalues_mxfp4_const[i]; - } - if (MmTypeA == GGML_TYPE_NVFP4) { - for (uint i = gl_LocalInvocationIndex.x; i < 128u; i += wgsize.x) { - ue4m3_fp32_lut[i] = ue4m3_to_fp32_build(i); - } - } - barrier(); - } -#endif } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp index 7540fd8711..40dbbcdc64 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp @@ -58,11 +58,6 @@ uint mm_load_vec_a() { case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_1: - case GGML_TYPE_IQ1_S: - case GGML_TYPE_IQ1_M: - case GGML_TYPE_IQ2_XXS: - case GGML_TYPE_IQ2_XS: - case GGML_TYPE_IQ2_S: return 8u; case GGML_TYPE_Q2_0: case GGML_TYPE_Q5_0: @@ -70,12 +65,6 @@ uint mm_load_vec_a() { case GGML_TYPE_Q2_K: case GGML_TYPE_Q4_K: case GGML_TYPE_Q5_K: - case GGML_TYPE_IQ3_XXS: - case GGML_TYPE_IQ3_S: - case GGML_TYPE_IQ4_XS: - case GGML_TYPE_IQ4_NL: - case GGML_TYPE_MXFP4: - case GGML_TYPE_NVFP4: return 4u; default: return 2u; @@ -114,31 +103,18 @@ layout (binding = 0) readonly buffer BUF_Q3_K { block_q3_K data[]; } a_q layout (binding = 0) readonly buffer BUF_Q4_K { block_q4_K data[]; } a_q4_k; layout (binding = 0) readonly buffer BUF_Q5_K { block_q5_K data[]; } a_q5_k; layout (binding = 0) readonly buffer BUF_Q6_K { block_q6_K data[]; } a_q6_k; -layout (binding = 0) readonly buffer BUF_IQ1_S { block_iq1_s data[]; } a_iq1_s; -layout (binding = 0) readonly buffer BUF_IQ1_M { block_iq1_m data[]; } a_iq1_m; -layout (binding = 0) readonly buffer BUF_IQ2_XXS { block_iq2_xxs data[]; } a_iq2_xxs; -layout (binding = 0) readonly buffer BUF_IQ2_XS { block_iq2_xs data[]; } a_iq2_xs; -layout (binding = 0) readonly buffer BUF_IQ2_S { block_iq2_s data[]; } a_iq2_s; -layout (binding = 0) readonly buffer BUF_IQ3_XXS { block_iq3_xxs data[]; } a_iq3_xxs; -layout (binding = 0) readonly buffer BUF_IQ3_S { block_iq3_s data[]; } a_iq3_s; -layout (binding = 0) readonly buffer BUF_IQ4_XS { block_iq4_xs data[]; } a_iq4_xs; -layout (binding = 0) readonly buffer BUF_MXFP4 { block_mxfp4 data[]; } a_mxfp4; -layout (binding = 0) readonly buffer BUF_NVFP4 { block_nvfp4 data[]; } a_nvfp4; // Packed16 aliases layout (binding = 0) readonly buffer BUF_Q4_0_P16 { block_q4_0_packed16 data[]; } a_q4_0_p16; layout (binding = 0) readonly buffer BUF_Q5_0_P16 { block_q5_0_packed16 data[]; } a_q5_0_p16; layout (binding = 0) readonly buffer BUF_Q8_0_P16 { block_q8_0_packed16 data[]; } a_q8_0_p16; layout (binding = 0) readonly buffer BUF_Q3_K_P16 { block_q3_K_packed16 data[]; } a_q3_k_p16; layout (binding = 0) readonly buffer BUF_Q6_K_P16 { block_q6_K_packed16 data[]; } a_q6_k_p16; -layout (binding = 0) readonly buffer BUF_IQ3_XXS_P16{ block_iq3_xxs_packed16 data[];} a_iq3_xxs_p16; -layout (binding = 0) readonly buffer BUF_IQ4_NL_P16 { block_iq4_nl_packed16 data[]; } a_iq4_nl_p16; // Packed32 aliases layout (binding = 0) readonly buffer BUF_Q4_1_P32 { block_q4_1_packed32 data[]; } a_q4_1_p32; layout (binding = 0) readonly buffer BUF_Q5_1_P32 { block_q5_1_packed32 data[]; } a_q5_1_p32; layout (binding = 0) readonly buffer BUF_Q2_K_P32 { block_q2_K_packed32 data[]; } a_q2_k_p32; layout (binding = 0) readonly buffer BUF_Q4_K_P32 { block_q4_K_packed32 data[]; } a_q4_k_p32; layout (binding = 0) readonly buffer BUF_Q5_K_P32 { block_q5_K_packed32 data[]; } a_q5_k_p32; -layout (binding = 0) readonly buffer BUF_IQ4_XS_P32 { block_iq4_xs_packed32 data[]; } a_iq4_xs_p32; #endif layout (binding = 1) readonly buffer B {B_TYPE data_b[];}; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp index b31e929d87..35ba6d7b46 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp @@ -50,14 +50,13 @@ layout (constant_id = 6) const uint subgroup_size = 32; uint mm_quant_k() { switch (MmTypeA) { case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: case GGML_TYPE_Q5_1: - case GGML_TYPE_Q8_0: case GGML_TYPE_IQ4_NL: case GGML_TYPE_MXFP4: + case GGML_TYPE_Q8_0: return 32u; case GGML_TYPE_Q1_0: return 128u; case GGML_TYPE_Q2_0: - case GGML_TYPE_NVFP4: return 64u; - default: // K-quants and IQ types + default: return 256u; } } @@ -113,17 +112,6 @@ layout (binding = 0) readonly buffer BUF_Q3_K { block_q3_K data[]; } a_q3_ layout (binding = 0) readonly buffer BUF_Q4_K { block_q4_K data[]; } a_q4_k; layout (binding = 0) readonly buffer BUF_Q5_K { block_q5_K data[]; } a_q5_k; layout (binding = 0) readonly buffer BUF_Q6_K { block_q6_K data[]; } a_q6_k; -layout (binding = 0) readonly buffer BUF_IQ1_S { block_iq1_s data[]; } a_iq1_s; -layout (binding = 0) readonly buffer BUF_IQ1_M { block_iq1_m data[]; } a_iq1_m; -layout (binding = 0) readonly buffer BUF_IQ2_XXS { block_iq2_xxs data[]; } a_iq2_xxs; -layout (binding = 0) readonly buffer BUF_IQ2_XS { block_iq2_xs data[]; } a_iq2_xs; -layout (binding = 0) readonly buffer BUF_IQ2_S { block_iq2_s data[]; } a_iq2_s; -layout (binding = 0) readonly buffer BUF_IQ3_XXS { block_iq3_xxs data[]; } a_iq3_xxs; -layout (binding = 0) readonly buffer BUF_IQ3_S { block_iq3_s data[]; } a_iq3_s; -layout (binding = 0) readonly buffer BUF_IQ4_XS { block_iq4_xs data[]; } a_iq4_xs; -layout (binding = 0) readonly buffer BUF_IQ4_NL { block_iq4_nl data[]; } a_iq4_nl; -layout (binding = 0) readonly buffer BUF_MXFP4 { block_mxfp4 data[]; } a_mxfp4; -layout (binding = 0) readonly buffer BUF_NVFP4 { block_nvfp4 data[]; } a_nvfp4; #endif layout (binding = 1) readonly buffer B {B_TYPE data_b[];}; layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; @@ -158,18 +146,7 @@ layout (binding = 1) readonly buffer B4 {B_TYPEV4 data_b_v4[];}; else if (MmTypeA == GGML_TYPE_Q3_K) { coopMatLoadTensorNV(mat, a_q3_k.data, pos_a, tensor_layout, dequantFuncQ3_K, dequantFuncQ3_K_v); } \ else if (MmTypeA == GGML_TYPE_Q4_K) { coopMatLoadTensorNV(mat, a_q4_k.data, pos_a, tensor_layout, dequantFuncQ4_K, dequantFuncQ4_K_v); } \ else if (MmTypeA == GGML_TYPE_Q5_K) { coopMatLoadTensorNV(mat, a_q5_k.data, pos_a, tensor_layout, dequantFuncQ5_K, dequantFuncQ5_K_v); } \ - else if (MmTypeA == GGML_TYPE_Q6_K) { coopMatLoadTensorNV(mat, a_q6_k.data, pos_a, tensor_layout, dequantFuncQ6_K, dequantFuncQ6_K_v); } \ - else if (MmTypeA == GGML_TYPE_IQ1_S) { coopMatLoadTensorNV(mat, a_iq1_s.data, pos_a, tensor_layout, dequantFuncIQ1_S, dequantFuncIQ1_S_v); } \ - else if (MmTypeA == GGML_TYPE_IQ1_M) { coopMatLoadTensorNV(mat, a_iq1_m.data, pos_a, tensor_layout, dequantFuncIQ1_M, dequantFuncIQ1_M_v); } \ - else if (MmTypeA == GGML_TYPE_IQ2_XXS) { coopMatLoadTensorNV(mat, a_iq2_xxs.data, pos_a, tensor_layout, dequantFuncIQ2_XXS, dequantFuncIQ2_XXS_v); } \ - else if (MmTypeA == GGML_TYPE_IQ2_XS) { coopMatLoadTensorNV(mat, a_iq2_xs.data, pos_a, tensor_layout, dequantFuncIQ2_XS, dequantFuncIQ2_XS_v); } \ - else if (MmTypeA == GGML_TYPE_IQ2_S) { coopMatLoadTensorNV(mat, a_iq2_s.data, pos_a, tensor_layout, dequantFuncIQ2_S, dequantFuncIQ2_S_v); } \ - else if (MmTypeA == GGML_TYPE_IQ3_XXS) { coopMatLoadTensorNV(mat, a_iq3_xxs.data, pos_a, tensor_layout, dequantFuncIQ3_XXS, dequantFuncIQ3_XXS_v); } \ - else if (MmTypeA == GGML_TYPE_IQ3_S) { coopMatLoadTensorNV(mat, a_iq3_s.data, pos_a, tensor_layout, dequantFuncIQ3_S, dequantFuncIQ3_S_v); } \ - else if (MmTypeA == GGML_TYPE_IQ4_XS) { coopMatLoadTensorNV(mat, a_iq4_xs.data, pos_a, tensor_layout, dequantFuncIQ4_XS, dequantFuncIQ4_XS_v); } \ - else if (MmTypeA == GGML_TYPE_IQ4_NL) { coopMatLoadTensorNV(mat, a_iq4_nl.data, pos_a, tensor_layout, dequantFuncIQ4_NL, dequantFuncIQ4_NL_v); } \ - else if (MmTypeA == GGML_TYPE_MXFP4) { coopMatLoadTensorNV(mat, a_mxfp4.data, pos_a, tensor_layout, dequantFuncMXFP4, dequantFuncMXFP4_v); } \ - else if (MmTypeA == GGML_TYPE_NVFP4) { coopMatLoadTensorNV(mat, a_nvfp4.data, pos_a, tensor_layout, dequantFuncNVFP4, dequantFuncNVFP4_v); } + else if (MmTypeA == GGML_TYPE_Q6_K) { coopMatLoadTensorNV(mat, a_q6_k.data, pos_a, tensor_layout, dequantFuncQ6_K, dequantFuncQ6_K_v); } #else #define COOPMAT_LOAD_A(mat, tensor_layout) \ if (MmTypeA == GGML_TYPE_Q1_0) { coopMatLoadTensorNV(mat, a_q1_0.data, pos_a, tensor_layout, dequantFuncQ1_0); } \ @@ -184,18 +161,7 @@ layout (binding = 1) readonly buffer B4 {B_TYPEV4 data_b_v4[];}; else if (MmTypeA == GGML_TYPE_Q3_K) { coopMatLoadTensorNV(mat, a_q3_k.data, pos_a, tensor_layout, dequantFuncQ3_K); } \ else if (MmTypeA == GGML_TYPE_Q4_K) { coopMatLoadTensorNV(mat, a_q4_k.data, pos_a, tensor_layout, dequantFuncQ4_K); } \ else if (MmTypeA == GGML_TYPE_Q5_K) { coopMatLoadTensorNV(mat, a_q5_k.data, pos_a, tensor_layout, dequantFuncQ5_K); } \ - else if (MmTypeA == GGML_TYPE_Q6_K) { coopMatLoadTensorNV(mat, a_q6_k.data, pos_a, tensor_layout, dequantFuncQ6_K); } \ - else if (MmTypeA == GGML_TYPE_IQ1_S) { coopMatLoadTensorNV(mat, a_iq1_s.data, pos_a, tensor_layout, dequantFuncIQ1_S); } \ - else if (MmTypeA == GGML_TYPE_IQ1_M) { coopMatLoadTensorNV(mat, a_iq1_m.data, pos_a, tensor_layout, dequantFuncIQ1_M); } \ - else if (MmTypeA == GGML_TYPE_IQ2_XXS) { coopMatLoadTensorNV(mat, a_iq2_xxs.data, pos_a, tensor_layout, dequantFuncIQ2_XXS); } \ - else if (MmTypeA == GGML_TYPE_IQ2_XS) { coopMatLoadTensorNV(mat, a_iq2_xs.data, pos_a, tensor_layout, dequantFuncIQ2_XS); } \ - else if (MmTypeA == GGML_TYPE_IQ2_S) { coopMatLoadTensorNV(mat, a_iq2_s.data, pos_a, tensor_layout, dequantFuncIQ2_S); } \ - else if (MmTypeA == GGML_TYPE_IQ3_XXS) { coopMatLoadTensorNV(mat, a_iq3_xxs.data, pos_a, tensor_layout, dequantFuncIQ3_XXS); } \ - else if (MmTypeA == GGML_TYPE_IQ3_S) { coopMatLoadTensorNV(mat, a_iq3_s.data, pos_a, tensor_layout, dequantFuncIQ3_S); } \ - else if (MmTypeA == GGML_TYPE_IQ4_XS) { coopMatLoadTensorNV(mat, a_iq4_xs.data, pos_a, tensor_layout, dequantFuncIQ4_XS); } \ - else if (MmTypeA == GGML_TYPE_IQ4_NL) { coopMatLoadTensorNV(mat, a_iq4_nl.data, pos_a, tensor_layout, dequantFuncIQ4_NL); } \ - else if (MmTypeA == GGML_TYPE_MXFP4) { coopMatLoadTensorNV(mat, a_mxfp4.data, pos_a, tensor_layout, dequantFuncMXFP4); } \ - else if (MmTypeA == GGML_TYPE_NVFP4) { coopMatLoadTensorNV(mat, a_nvfp4.data, pos_a, tensor_layout, dequantFuncNVFP4); } + else if (MmTypeA == GGML_TYPE_Q6_K) { coopMatLoadTensorNV(mat, a_q6_k.data, pos_a, tensor_layout, dequantFuncQ6_K); } #endif #endif #else diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl index c29219d3f4..2c5d69c1fd 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl @@ -69,6 +69,276 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin } else { store_a(col, row, FLOAT_TYPEV2(0.0f)); } +#elif defined(DATA_A_IQ1_S) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 32; + const uint ib32 = (idx % 32) / 4; + const uint ib8 = idx % 32; + + const float d = float(data_a[ib].d); + const uint qh = data_a[ib].qh[ib32]; + const uint qs = data_a[ib].qs[ib8]; + const float dl = d * (2 * bitfieldExtract(qh, 12, 3) + 1); + const float delta = ((qh & 0x8000) != 0) ? -IQ1S_DELTA : IQ1S_DELTA; + const int16_t grid = int16_t(iq1s_grid[qs | (bitfieldExtract(qh, 3 * int(ib8 & 3), 3) << 8)]); + + [[unroll]] for (int k = 0; k < 4; ++k) { + store_a(col, k_pair + k, FLOAT_TYPEV2(dl * (bitfieldExtract(grid, 4 * k , 2) + delta), + dl * (bitfieldExtract(grid, 4 * k + 2, 2) + delta))); + + } +#elif defined(DATA_A_IQ1_M) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 32; + const uint ib8 = idx % 32; + const uint ib16 = ib8 / 2; + + const uint16_t[4] scales = data_a[ib].scales; + const u16vec4 s = u16vec4(scales[0], scales[1], scales[2], scales[3]) >> 12; + const float d = float(unpackHalf2x16(s.x | (s.y << 4) | (s.z << 8) | (s.w << 12)).x); + const uint sc = scales[ib8 / 8]; + const uint qs = data_a[ib].qs[ib8]; + const uint qh = data_a[ib].qh[ib16] >> (4 * (ib8 & 1)); + const float dl = d * (2 * bitfieldExtract(sc, 3 * int(ib16 & 3), 3) + 1); + const float delta = ((qh & 8) != 0) ? -IQ1M_DELTA : IQ1M_DELTA; + const int16_t grid = int16_t(iq1s_grid[qs | ((qh & 7) << 8)]); + + [[unroll]] for (int k = 0; k < 4; ++k) { + store_a(col, k_pair + k, FLOAT_TYPEV2(dl * (bitfieldExtract(grid, 4 * k , 2) + delta), + dl * (bitfieldExtract(grid, 4 * k + 2, 2) + delta))); + + } +#elif defined(DATA_A_IQ2_XXS) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 32; + const uint ib32 = (idx % 32) / 4; + const uint ib8 = idx % 4; + + const float d = float(data_a[ib].d); + const uint qs = data_a[ib].qs[8 * ib32 + ib8]; + const uint signs = pack32(u8vec4( + data_a[ib].qs[8*ib32 + 4], + data_a[ib].qs[8*ib32 + 5], + data_a[ib].qs[8*ib32 + 6], + data_a[ib].qs[8*ib32 + 7] + )); + const FLOAT_TYPE db = FLOAT_TYPE(d * 0.25 * (0.5 + (signs >> 28))); + const uint32_t sign7 = bitfieldExtract(signs, 7 * int(ib8), 7); + const uint sign = sign7 | (bitCount(sign7) << 7); + const uvec2 grid = iq2_grid[qs]; + const vec4 grid0 = vec4(unpack8(grid.x)); + const vec4 grid1 = vec4(unpack8(grid.y)); + + store_a(col, k_pair, db * FLOAT_TYPEV2((sign & 1) != 0 ? -grid0.x : grid0.x, + (sign & 2) != 0 ? -grid0.y : grid0.y)); + + store_a(col, k_pair + 1, db * FLOAT_TYPEV2((sign & 4) != 0 ? -grid0.z : grid0.z, + (sign & 8) != 0 ? -grid0.w : grid0.w)); + + store_a(col, k_pair + 2, db * FLOAT_TYPEV2((sign & 16) != 0 ? -grid1.x : grid1.x, + (sign & 32) != 0 ? -grid1.y : grid1.y)); + + store_a(col, k_pair + 3, db * FLOAT_TYPEV2((sign & 64) != 0 ? -grid1.z : grid1.z, + (sign & 128) != 0 ? -grid1.w : grid1.w)); + +#elif defined(DATA_A_IQ2_XS) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 32; + const uint ib32 = (idx % 32) / 4; + const uint ib8 = idx % 4; + + const float d = float(data_a[ib].d); + const uint scale = (data_a[ib].scales[ib32] >> (2 * (ib8 & 2))) & 0xf; + const FLOAT_TYPE db = FLOAT_TYPE(d * 0.25 * (0.5 + scale)); + const uint qs = data_a[ib].qs[4 * ib32 + ib8]; + const uint sign7 = qs >> 9; + const uint sign = sign7 | (bitCount(sign7) << 7); + const uvec2 grid = iq2_grid[qs & 511]; + const vec4 grid0 = vec4(unpack8(grid.x)); + const vec4 grid1 = vec4(unpack8(grid.y)); + + store_a(col, k_pair, db * FLOAT_TYPEV2((sign & 1) != 0 ? -grid0.x : grid0.x, + (sign & 2) != 0 ? -grid0.y : grid0.y)); + + store_a(col, k_pair + 1, db * FLOAT_TYPEV2((sign & 4) != 0 ? -grid0.z : grid0.z, + (sign & 8) != 0 ? -grid0.w : grid0.w)); + + store_a(col, k_pair + 2, db * FLOAT_TYPEV2((sign & 16) != 0 ? -grid1.x : grid1.x, + (sign & 32) != 0 ? -grid1.y : grid1.y)); + + store_a(col, k_pair + 3, db * FLOAT_TYPEV2((sign & 64) != 0 ? -grid1.z : grid1.z, + (sign & 128) != 0 ? -grid1.w : grid1.w)); + +#elif defined(DATA_A_IQ2_S) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 32; + const uint ib8 = idx % 32; + const uint ib32 = ib8 / 4; + + const uint scale = (data_a[ib].scales[ib32] >> (2 * (ib8 & 2))) & 0xf; + const uint qs = data_a[ib].qs[ib8]; + const uint qh = data_a[ib].qh[ib32]; + const uint qhshift = 2 * (ib8 % 4); + const uint sign = data_a[ib].qs[QUANT_K_IQ2_S / 8 + ib8]; + + const float d = float(data_a[ib].d); + const FLOAT_TYPE db = FLOAT_TYPE(d * 0.25 * (0.5 + scale)); + const uvec2 grid = iq2_grid[qs | ((qh << (8 - qhshift)) & 0x300)]; + const vec4 grid0 = vec4(unpack8(grid.x)); + const vec4 grid1 = vec4(unpack8(grid.y)); + + store_a(col, k_pair, db * FLOAT_TYPEV2((sign & 1) != 0 ? -grid0.x : grid0.x, + (sign & 2) != 0 ? -grid0.y : grid0.y)); + + store_a(col, k_pair + 1, db * FLOAT_TYPEV2((sign & 4) != 0 ? -grid0.z : grid0.z, + (sign & 8) != 0 ? -grid0.w : grid0.w)); + + store_a(col, k_pair + 2, db * FLOAT_TYPEV2((sign & 16) != 0 ? -grid1.x : grid1.x, + (sign & 32) != 0 ? -grid1.y : grid1.y)); + + store_a(col, k_pair + 3, db * FLOAT_TYPEV2((sign & 64) != 0 ? -grid1.z : grid1.z, + (sign & 128) != 0 ? -grid1.w : grid1.w)); + +#elif defined(DATA_A_IQ3_XXS) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 64; + const uint iqs = idx % 64; + const uint is = QUANT_K_IQ3_XXS / 4 + 4 * (iqs / 8); + + const float d = float(data_a[ib].d); + const uint qs = data_a[ib].qs[iqs]; + const uint signs = pack32(u16vec2( + data_a_packed16[ib].qs[is/2], + data_a_packed16[ib].qs[is/2+1] + )); + const float db = d * 0.5 * (0.5 + (signs >> 28)); + const uint32_t sign7 = bitfieldExtract(signs, 7 * (int(iqs / 2) % 4), 7); + const uint sign = (sign7 | (bitCount(sign7) << 7)) >> (4 * (idx % 2)); + const uint grid = iq3_grid[qs]; + const vec4 v = db * vec4(unpack8(grid)); + + store_a(col, k_pair, FLOAT_TYPEV2((sign & 1) != 0 ? -v.x : v.x, + (sign & 2) != 0 ? -v.y : v.y)); + + store_a(col, k_pair + 1, FLOAT_TYPEV2((sign & 4) != 0 ? -v.z : v.z, + (sign & 8) != 0 ? -v.w : v.w)); + +#elif defined(DATA_A_IQ3_S) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 64; + const uint iqs = idx % 64; + const uint iqh = iqs / 8; + + const float d = float(data_a[ib].d); + const uint qs = data_a[ib].qs[iqs]; + const uint qh = data_a[ib].qh[iqh]; + const int8_t sign = int8_t(data_a[ib].signs[iqs / 2] >> (4 * (idx % 2))); + const uint scale = data_a[ib].scales[iqs / 16]; + const i8vec2 sign01 = i8vec2(1 - (2 & i8vec2(sign << 1, sign))); + const float db = d * (1 + 2 * ((scale >> (4 * (iqh & 1))) & 0xf)); + const uint32_t grid = iq3_grid[qs | ((qh << (8 - (iqs % 8))) & 256)]; + const vec4 v = db * vec4(unpack8(grid)); + + store_a(col, k_pair, FLOAT_TYPEV2((sign & 1) != 0 ? -v.x : v.x, + (sign & 2) != 0 ? -v.y : v.y)); + + store_a(col, k_pair + 1, FLOAT_TYPEV2((sign & 4) != 0 ? -v.z : v.z, + (sign & 8) != 0 ? -v.w : v.w)); + +#elif defined(DATA_A_IQ4_XS) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 2; + + const uint ib = idx / 64; + const uint ib32 = (idx % 64) / 8; + const uint iq = 4 * ib32 + (idx % 4); + + const uint sl = (data_a[ib].scales_l[ib32/2] >> (4 * (ib32 & 1))) & 0xF; + const uint sh = ((data_a[ib].scales_h) >> (2 * ib32)) & 3; + const uint qshift = idx & 4; + u8vec4 qs = unpack8((uint(data_a_packed32[ib].qs[iq]) >> qshift) & 0x0F0F0F0F); + + const float d = float(data_a[ib].d); + const vec4 v = d * float(int(sl | (sh << 4)) - 32) * vec4(kvalues_iq4nl[qs.x], kvalues_iq4nl[qs.y], kvalues_iq4nl[qs.z], kvalues_iq4nl[qs.w]); + + store_a(col, k_pair, FLOAT_TYPEV2(v.xy)); + store_a(col, k_pair + 1, FLOAT_TYPEV2(v.zw)); +#elif defined(DATA_A_IQ4_NL) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 4; + + const uint ib = idx / 8; + const uint iqs = idx & 0x07; + + const FLOAT_TYPE d = FLOAT_TYPE(data_a_packed16[ib].d); + const uint vui = uint(data_a_packed16[ib].qs[iqs]); + + store_a(col, k_pair, d * FLOAT_TYPEV2(kvalues_iq4nl[vui & 0xF], + kvalues_iq4nl[bitfieldExtract(vui, 8, 4)])); + + store_a(col, k_pair + 8, d * FLOAT_TYPEV2(kvalues_iq4nl[bitfieldExtract(vui, 4, 4)], + kvalues_iq4nl[vui >> 12])); + +#elif defined(DATA_A_MXFP4) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint k_pair = row * LOAD_VEC_A / 4; + + const uint ib = idx / 8; + const uint iqs = (idx & 0x07) * 2; + + const uint vui = uint(data_a[ib].qs[iqs]); + const uint vui2 = uint(data_a[ib].qs[iqs+1]); + +#ifdef USE_OCP_FP4 + const float d = e8m0_to_fp32(data_a[ib].e); + const u8vec2 packed = u8vec2(vui, vui2); + store_a(col, k_pair, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 0u)) * FLOAT_TYPE(d)); + store_a(col, k_pair + 8, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 4u)) * FLOAT_TYPE(d)); +#else + const float d = e8m0_to_fp32(data_a[ib].e) * 0.5; + store_a(col, k_pair, FLOAT_TYPEV2(kvalues_mxfp4[vui & 0xF] * d, + kvalues_mxfp4[vui2 & 0xF] * d)); + + store_a(col, k_pair + 8, FLOAT_TYPEV2(kvalues_mxfp4[vui >> 4] * d, + kvalues_mxfp4[vui2 >> 4] * d)); + +#endif +#elif defined(DATA_A_NVFP4) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + const uint eff_row = (row & 3) + (row & ~3) * 2; + + const uint ib = idx / 16u; + const uint sub = (idx & 0xC) >> 2; + const uint iqs = (idx & 0xF) * 2; + const uint vui = uint(data_a[ib].qs[iqs]); + const uint vui2 = uint(data_a[ib].qs[iqs+1]); + +#ifdef USE_OCP_FP4 + const FLOAT_TYPE d = FLOAT_TYPE(ue4m3_from_bits(data_a[ib].d[sub])); + const u8vec2 packed = u8vec2(vui, vui2); + store_a(col, eff_row, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 0u)) * d); + store_a(col, eff_row + 4, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 4u)) * d); +#else + const float d = ue4m3_to_fp32(data_a[ib].d[sub]) * 0.5; + store_a(col, eff_row, FLOAT_TYPEV2(kvalues_mxfp4[vui & 0xF] * d, + kvalues_mxfp4[vui2 & 0xF] * d)); + store_a(col, eff_row + 4, FLOAT_TYPEV2(kvalues_mxfp4[vui >> 4] * d, + kvalues_mxfp4[vui2 >> 4] * d)); +#endif #else if (MmTypeA == GGML_TYPE_Q4_0) { const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; @@ -338,279 +608,7 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin const vec2 q = (vec2(unpack8(ql | (qh << 4)).xy) - 32) * dscale; store_a(col, k_pair, FLOAT_TYPEV2(q.x, q.y)); - } else if (MmTypeA == GGML_TYPE_IQ1_S) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 32; // 8 values per idx - const uint ib32 = (idx % 32) / 4; // 0..7 - const uint ib8 = idx % 32; - - const float d = float(a_iq1_s.data[ib].d); - const uint qh = a_iq1_s.data[ib].qh[ib32]; - const uint qs = a_iq1_s.data[ib].qs[ib8]; - const float dl = d * (2 * bitfieldExtract(qh, 12, 3) + 1); - const float delta = ((qh & 0x8000) != 0) ? -IQ1S_DELTA : IQ1S_DELTA; - const int16_t grid = int16_t(iq1s_grid[qs | (bitfieldExtract(qh, 3 * int(ib8 & 3), 3) << 8)]); - - [[unroll]] for (int k = 0; k < 4; ++k) { - store_a(col, k_pair + k, FLOAT_TYPEV2(dl * (bitfieldExtract(grid, 4 * k , 2) + delta), - dl * (bitfieldExtract(grid, 4 * k + 2, 2) + delta))); - - } - } else if (MmTypeA == GGML_TYPE_IQ1_M) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 32; // 8 values per idx - const uint ib8 = idx % 32; - const uint ib16 = ib8 / 2; - - const uint16_t[4] scales = a_iq1_m.data[ib].scales; - const u16vec4 s = u16vec4(scales[0], scales[1], scales[2], scales[3]) >> 12; - const float d = float(unpackHalf2x16(s.x | (s.y << 4) | (s.z << 8) | (s.w << 12)).x); - const uint sc = scales[ib8 / 8]; - const uint qs = a_iq1_m.data[ib].qs[ib8]; - const uint qh = a_iq1_m.data[ib].qh[ib16] >> (4 * (ib8 & 1)); - const float dl = d * (2 * bitfieldExtract(sc, 3 * int(ib16 & 3), 3) + 1); - const float delta = ((qh & 8) != 0) ? -IQ1M_DELTA : IQ1M_DELTA; - const int16_t grid = int16_t(iq1s_grid[qs | ((qh & 7) << 8)]); - - [[unroll]] for (int k = 0; k < 4; ++k) { - store_a(col, k_pair + k, FLOAT_TYPEV2(dl * (bitfieldExtract(grid, 4 * k , 2) + delta), - dl * (bitfieldExtract(grid, 4 * k + 2, 2) + delta))); - - } - } else if (MmTypeA == GGML_TYPE_IQ2_XXS) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 32; // 8 values per idx - const uint ib32 = (idx % 32) / 4; // 0..7 - const uint ib8 = idx % 4; - - const float d = float(a_iq2_xxs.data[ib].d); - const uint qs = a_iq2_xxs.data[ib].qs[8 * ib32 + ib8]; - const uint signs = pack32(u8vec4( - a_iq2_xxs.data[ib].qs[8*ib32 + 4], - a_iq2_xxs.data[ib].qs[8*ib32 + 5], - a_iq2_xxs.data[ib].qs[8*ib32 + 6], - a_iq2_xxs.data[ib].qs[8*ib32 + 7] - )); - const FLOAT_TYPE db = FLOAT_TYPE(d * 0.25 * (0.5 + (signs >> 28))); - const uint32_t sign7 = bitfieldExtract(signs, 7 * int(ib8), 7); - const uint sign = sign7 | (bitCount(sign7) << 7); - const uvec2 grid = iq2_grid[qs]; - const vec4 grid0 = vec4(unpack8(grid.x)); - const vec4 grid1 = vec4(unpack8(grid.y)); - - store_a(col, k_pair, db * FLOAT_TYPEV2((sign & 1) != 0 ? -grid0.x : grid0.x, - (sign & 2) != 0 ? -grid0.y : grid0.y)); - - store_a(col, k_pair + 1, db * FLOAT_TYPEV2((sign & 4) != 0 ? -grid0.z : grid0.z, - (sign & 8) != 0 ? -grid0.w : grid0.w)); - - store_a(col, k_pair + 2, db * FLOAT_TYPEV2((sign & 16) != 0 ? -grid1.x : grid1.x, - (sign & 32) != 0 ? -grid1.y : grid1.y)); - - store_a(col, k_pair + 3, db * FLOAT_TYPEV2((sign & 64) != 0 ? -grid1.z : grid1.z, - (sign & 128) != 0 ? -grid1.w : grid1.w)); - - } else if (MmTypeA == GGML_TYPE_IQ2_XS) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 32; // 8 values per idx - const uint ib32 = (idx % 32) / 4; // 0..7 - const uint ib8 = idx % 4; // 0..3 - - const float d = float(a_iq2_xs.data[ib].d); - const uint scale = (a_iq2_xs.data[ib].scales[ib32] >> (2 * (ib8 & 2))) & 0xf; - const FLOAT_TYPE db = FLOAT_TYPE(d * 0.25 * (0.5 + scale)); - const uint qs = a_iq2_xs.data[ib].qs[4 * ib32 + ib8]; - const uint sign7 = qs >> 9; - const uint sign = sign7 | (bitCount(sign7) << 7); - const uvec2 grid = iq2_grid[qs & 511]; - const vec4 grid0 = vec4(unpack8(grid.x)); - const vec4 grid1 = vec4(unpack8(grid.y)); - - store_a(col, k_pair, db * FLOAT_TYPEV2((sign & 1) != 0 ? -grid0.x : grid0.x, - (sign & 2) != 0 ? -grid0.y : grid0.y)); - - store_a(col, k_pair + 1, db * FLOAT_TYPEV2((sign & 4) != 0 ? -grid0.z : grid0.z, - (sign & 8) != 0 ? -grid0.w : grid0.w)); - - store_a(col, k_pair + 2, db * FLOAT_TYPEV2((sign & 16) != 0 ? -grid1.x : grid1.x, - (sign & 32) != 0 ? -grid1.y : grid1.y)); - - store_a(col, k_pair + 3, db * FLOAT_TYPEV2((sign & 64) != 0 ? -grid1.z : grid1.z, - (sign & 128) != 0 ? -grid1.w : grid1.w)); - - } else if (MmTypeA == GGML_TYPE_IQ2_S) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 32; // 8 values per idx - const uint ib8 = idx % 32; // 0..31 - const uint ib32 = ib8 / 4; // 0..7 - - const uint scale = (a_iq2_s.data[ib].scales[ib32] >> (2 * (ib8 & 2))) & 0xf; - const uint qs = a_iq2_s.data[ib].qs[ib8]; - const uint qh = a_iq2_s.data[ib].qh[ib32]; - const uint qhshift = 2 * (ib8 % 4); - const uint sign = a_iq2_s.data[ib].qs[QUANT_K_IQ2_S / 8 + ib8]; - - const float d = float(a_iq2_s.data[ib].d); - const FLOAT_TYPE db = FLOAT_TYPE(d * 0.25 * (0.5 + scale)); - const uvec2 grid = iq2_grid[qs | ((qh << (8 - qhshift)) & 0x300)]; - const vec4 grid0 = vec4(unpack8(grid.x)); - const vec4 grid1 = vec4(unpack8(grid.y)); - - store_a(col, k_pair, db * FLOAT_TYPEV2((sign & 1) != 0 ? -grid0.x : grid0.x, - (sign & 2) != 0 ? -grid0.y : grid0.y)); - - store_a(col, k_pair + 1, db * FLOAT_TYPEV2((sign & 4) != 0 ? -grid0.z : grid0.z, - (sign & 8) != 0 ? -grid0.w : grid0.w)); - - store_a(col, k_pair + 2, db * FLOAT_TYPEV2((sign & 16) != 0 ? -grid1.x : grid1.x, - (sign & 32) != 0 ? -grid1.y : grid1.y)); - - store_a(col, k_pair + 3, db * FLOAT_TYPEV2((sign & 64) != 0 ? -grid1.z : grid1.z, - (sign & 128) != 0 ? -grid1.w : grid1.w)); - - } else if (MmTypeA == GGML_TYPE_IQ3_XXS) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 64; // 4 values per idx - const uint iqs = idx % 64; // 0..63 - const uint is = QUANT_K_IQ3_XXS / 4 + 4 * (iqs / 8); // 8 values - - const float d = float(a_iq3_xxs.data[ib].d); - const uint qs = a_iq3_xxs.data[ib].qs[iqs]; - const uint signs = pack32(u16vec2( - a_iq3_xxs_p16.data[ib].qs[is/2], - a_iq3_xxs_p16.data[ib].qs[is/2+1] - )); - const float db = d * 0.5 * (0.5 + (signs >> 28)); - const uint32_t sign7 = bitfieldExtract(signs, 7 * (int(iqs / 2) % 4), 7); - const uint sign = (sign7 | (bitCount(sign7) << 7)) >> (4 * (idx % 2)); - const uint grid = iq3_grid[qs]; - const vec4 v = db * vec4(unpack8(grid)); - - store_a(col, k_pair, FLOAT_TYPEV2((sign & 1) != 0 ? -v.x : v.x, - (sign & 2) != 0 ? -v.y : v.y)); - - store_a(col, k_pair + 1, FLOAT_TYPEV2((sign & 4) != 0 ? -v.z : v.z, - (sign & 8) != 0 ? -v.w : v.w)); - - } else if (MmTypeA == GGML_TYPE_IQ3_S) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 64; // 4 values per idx - const uint iqs = idx % 64; // 0..63 - const uint iqh = iqs / 8; - - const float d = float(a_iq3_s.data[ib].d); - const uint qs = a_iq3_s.data[ib].qs[iqs]; - const uint qh = a_iq3_s.data[ib].qh[iqh]; - const int8_t sign = int8_t(a_iq3_s.data[ib].signs[iqs / 2] >> (4 * (idx % 2))); - const uint scale = a_iq3_s.data[ib].scales[iqs / 16]; - const i8vec2 sign01 = i8vec2(1 - (2 & i8vec2(sign << 1, sign))); - const float db = d * (1 + 2 * ((scale >> (4 * (iqh & 1))) & 0xf)); - const uint32_t grid = iq3_grid[qs | ((qh << (8 - (iqs % 8))) & 256)]; - const vec4 v = db * vec4(unpack8(grid)); - - store_a(col, k_pair, FLOAT_TYPEV2((sign & 1) != 0 ? -v.x : v.x, - (sign & 2) != 0 ? -v.y : v.y)); - - store_a(col, k_pair + 1, FLOAT_TYPEV2((sign & 4) != 0 ? -v.z : v.z, - (sign & 8) != 0 ? -v.w : v.w)); - - } else if (MmTypeA == GGML_TYPE_IQ4_XS) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 2; - - const uint ib = idx / 64; // 4 values per idx - const uint ib32 = (idx % 64) / 8; // 0..7 - const uint iq = 4 * ib32 + (idx % 4); - - const uint sl = (a_iq4_xs.data[ib].scales_l[ib32/2] >> (4 * (ib32 & 1))) & 0xF; - const uint sh = ((a_iq4_xs.data[ib].scales_h) >> (2 * ib32)) & 3; - const uint qshift = idx & 4; - u8vec4 qs = unpack8((uint(a_iq4_xs_p32.data[ib].qs[iq]) >> qshift) & 0x0F0F0F0F); - - const float d = float(a_iq4_xs.data[ib].d); - const vec4 v = d * float(int(sl | (sh << 4)) - 32) * vec4(kvalues_iq4nl[qs.x], kvalues_iq4nl[qs.y], kvalues_iq4nl[qs.z], kvalues_iq4nl[qs.w]); - - store_a(col, k_pair, FLOAT_TYPEV2(v.xy)); - store_a(col, k_pair + 1, FLOAT_TYPEV2(v.zw)); - } else if (MmTypeA == GGML_TYPE_IQ4_NL) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 4; - - const uint ib = idx / 8; - const uint iqs = idx & 0x07; - - const FLOAT_TYPE d = FLOAT_TYPE(a_iq4_nl_p16.data[ib].d); - const uint vui = uint(a_iq4_nl_p16.data[ib].qs[iqs]); - - store_a(col, k_pair, d * FLOAT_TYPEV2(kvalues_iq4nl[vui & 0xF], - kvalues_iq4nl[bitfieldExtract(vui, 8, 4)])); - - store_a(col, k_pair + 8, d * FLOAT_TYPEV2(kvalues_iq4nl[bitfieldExtract(vui, 4, 4)], - kvalues_iq4nl[vui >> 12])); - - } else if (MmTypeA == GGML_TYPE_MXFP4) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - const uint k_pair = row * mm_load_vec_a() / 4; - - const uint ib = idx / 8; - const uint iqs = (idx & 0x07) * 2; - - const uint vui = uint(a_mxfp4.data[ib].qs[iqs]); - const uint vui2 = uint(a_mxfp4.data[ib].qs[iqs+1]); - -#ifdef USE_OCP_FP4 - const float d = e8m0_to_fp32(a_mxfp4.data[ib].e); - const u8vec2 packed = u8vec2(vui, vui2); - store_a(col, k_pair, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 0u)) * FLOAT_TYPE(d)); - store_a(col, k_pair + 8, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 4u)) * FLOAT_TYPE(d)); -#else - const float d = e8m0_to_fp32(a_mxfp4.data[ib].e) * 0.5; - store_a(col, k_pair, FLOAT_TYPEV2(kvalues_mxfp4[vui & 0xF] * d, - kvalues_mxfp4[vui2 & 0xF] * d)); - - store_a(col, k_pair + 8, FLOAT_TYPEV2(kvalues_mxfp4[vui >> 4] * d, - kvalues_mxfp4[vui2 >> 4] * d)); - -#endif - } else if (MmTypeA == GGML_TYPE_NVFP4) { - const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row; - // lo and hi nibbles are 8 elements apart, which doesn't quite line up with - // how the thread mapping and buf_idx calculation works for other types. - const uint eff_row = (row & 3) + (row & ~3) * 2; - - const uint ib = idx / 16u; - const uint sub = (idx & 0xC) >> 2; - const uint iqs = (idx & 0xF) * 2; - const uint vui = uint(a_nvfp4.data[ib].qs[iqs]); - const uint vui2 = uint(a_nvfp4.data[ib].qs[iqs+1]); - -#ifdef USE_OCP_FP4 - const FLOAT_TYPE d = FLOAT_TYPE(ue4m3_from_bits(a_nvfp4.data[ib].d[sub])); - const u8vec2 packed = u8vec2(vui, vui2); - store_a(col, eff_row, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 0u)) * d); - store_a(col, eff_row + 4, FLOAT_TYPEV2(bitcastExtractfe2m1EXT(packed, 4u)) * d); -#else - const float d = ue4m3_to_fp32(a_nvfp4.data[ib].d[sub]) * 0.5; - store_a(col, eff_row, FLOAT_TYPEV2(kvalues_mxfp4[vui & 0xF] * d, - kvalues_mxfp4[vui2 & 0xF] * d)); - store_a(col, eff_row + 4, FLOAT_TYPEV2(kvalues_mxfp4[vui >> 4] * d, - kvalues_mxfp4[vui2 >> 4] * d)); -#endif -} + } #endif } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl index b5936a3493..c50f14426c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl @@ -515,7 +515,7 @@ struct block_iq1_m_packed64 { #define A_TYPE_PACKED32 block_iq1_m_packed32 #endif -#if defined(DATA_A_IQ1_S) || defined(DATA_A_IQ1_M) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ1_S) || defined(DATA_A_IQ1_M) #define IQ1S_DELTA 0.125f #define IQ1M_DELTA 0.125f @@ -915,17 +915,10 @@ const uint32_t[2048] iq1s_grid_gpu_const = { }; #endif -#ifdef MULMAT_QUANT -shared uint16_t iq1s_grid[(MmTypeA == GGML_TYPE_IQ1_S || MmTypeA == GGML_TYPE_IQ1_M) ? 2048 : 8]; -#if defined(NEEDS_IQ1S_GRID_GPU) -shared uint32_t iq1s_grid_gpu[(MmTypeA == GGML_TYPE_IQ1_S || MmTypeA == GGML_TYPE_IQ1_M) ? 2048 : 8]; -#endif -#else shared uint16_t iq1s_grid[2048]; #if defined(NEEDS_IQ1S_GRID_GPU) shared uint32_t iq1s_grid_gpu[2048]; #endif -#endif #if defined(DATA_A_IQ1_S) || defined(DATA_A_IQ1_M) #define NEEDS_INIT_IQ_SHMEM @@ -953,10 +946,8 @@ void init_iq_shmem(uvec3 wgsize) #endif #endif -#if defined(DATA_A_IQ2_XXS) || defined(DATA_A_IQ2_XS) || defined(DATA_A_IQ2_S) || defined(MULMAT_QUANT) -#ifdef MULMAT_QUANT -shared uvec2 iq2_grid[MmTypeA == GGML_TYPE_IQ2_S ? 1024 : MmTypeA == GGML_TYPE_IQ2_XS ? 512 : MmTypeA == GGML_TYPE_IQ2_XXS ? 256 : 8]; -#elif defined(DATA_A_IQ2_S) +#if defined(DATA_A_IQ2_XXS) || defined(DATA_A_IQ2_XS) || defined(DATA_A_IQ2_S) +#if defined(DATA_A_IQ2_S) shared uvec2 iq2_grid[1024]; #elif defined(DATA_A_IQ2_XS) shared uvec2 iq2_grid[512]; @@ -980,7 +971,7 @@ struct block_iq2_xxs_packed16 uint16_t qs[QUANT_K_IQ2_XXS/8]; }; -#if defined(DATA_A_IQ2_XXS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ2_XXS) const uvec2[256] iq2xxs_grid_const = { uvec2(0x08080808, 0x08080808), uvec2(0x0808082b, 0x08080808), uvec2(0x08081919, 0x08080808), uvec2(0x08082b08, 0x08080808), @@ -1088,7 +1079,7 @@ struct block_iq2_xs_packed16 uint16_t scales[QUANT_K_IQ2_XS/64]; }; -#if defined(DATA_A_IQ2_XS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ2_XS) const uvec2 iq2xs_grid_const[512] = { uvec2(0x08080808, 0x08080808), uvec2(0x0808082b, 0x08080808), uvec2(0x08081919, 0x08080808), uvec2(0x08082b08, 0x08080808), @@ -1262,7 +1253,7 @@ struct block_iq2_s_packed16 uint16_t scales[QUANT_K_IQ2_S/64]; }; -#if defined(DATA_A_IQ2_S) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ2_S) const uvec2 iq2s_grid_const[1024] = { uvec2(0x08080808, 0x08080808), uvec2(0x0808082b, 0x08080808), uvec2(0x08081919, 0x08080808), uvec2(0x08082b08, 0x08080808), @@ -1545,10 +1536,8 @@ void init_iq_shmem(uvec3 wgsize) #endif #endif -#if defined(DATA_A_IQ3_XXS) || defined(DATA_A_IQ3_S) || defined(MULMAT_QUANT) -#ifdef MULMAT_QUANT -shared uint32_t iq3_grid[MmTypeA == GGML_TYPE_IQ3_S ? 512 : MmTypeA == GGML_TYPE_IQ3_XXS ? 256 : 8]; -#elif defined(DATA_A_IQ3_S) +#if defined(DATA_A_IQ3_XXS) || defined(DATA_A_IQ3_S) +#if defined(DATA_A_IQ3_S) shared uint32_t iq3_grid[512]; #else shared uint32_t iq3_grid[256]; @@ -1570,7 +1559,7 @@ struct block_iq3_xxs_packed16 uint16_t qs[QUANT_K_IQ3_XXS/8 + QUANT_K_IQ3_XXS/16]; }; -#if defined(DATA_A_IQ3_XXS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ3_XXS) const uint32_t iq3xxs_grid_const[256] = { 0x04040404, 0x04040414, 0x04040424, 0x04040c0c, 0x04040c1c, 0x04040c3e, 0x04041404, 0x04041414, @@ -1650,7 +1639,7 @@ struct block_iq3_s_packed16 uint16_t scales[QUANT_K_IQ3_S/64/2]; }; -#if defined(DATA_A_IQ3_S) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ3_S) const uint32_t iq3s_grid_const[512] = { 0x01010101, 0x01010103, 0x01010105, 0x0101010b, 0x0101010f, 0x01010301, 0x01010303, 0x01010305, @@ -1845,17 +1834,13 @@ struct block_nvfp4_packed32 #define A_TYPE_PACKED32 block_nvfp4_packed32 #endif -#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS) || defined(MULMAT_QUANT) +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS) const int8_t kvalues_iq4nl_const[16] = { int8_t(-127), int8_t(-104), int8_t(-83), int8_t(-65), int8_t(-49), int8_t(-35), int8_t(-22), int8_t(-10), int8_t(1), int8_t(13), int8_t(25), int8_t(38), int8_t(53), int8_t(69), int8_t(89), int8_t(113) }; -#ifdef MULMAT_QUANT -shared FLOAT_TYPE kvalues_iq4nl[(MmTypeA == GGML_TYPE_IQ4_NL || MmTypeA == GGML_TYPE_IQ4_XS) ? 16 : 8]; -#else shared FLOAT_TYPE kvalues_iq4nl[16]; -#endif #if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS) #define NEEDS_INIT_IQ_SHMEM @@ -1870,26 +1855,18 @@ void init_iq_shmem(uvec3 wgsize) #endif #endif -#if defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) || defined(MULMAT_QUANT) +#if defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) #if !defined(USE_OCP_FP4) const int8_t kvalues_mxfp4_const[16] = { int8_t(0), int8_t(1), int8_t(2), int8_t(3), int8_t(4), int8_t(6), int8_t(8), int8_t(12), int8_t(0), int8_t(-1), int8_t(-2), int8_t(-3), int8_t(-4), int8_t(-6), int8_t(-8), int8_t(-12), }; -#ifdef MULMAT_QUANT -shared int8_t kvalues_mxfp4[(MmTypeA == GGML_TYPE_MXFP4 || MmTypeA == GGML_TYPE_NVFP4) ? 16 : 8]; -#else shared int8_t kvalues_mxfp4[16]; #endif -#endif -#if (defined(DATA_A_NVFP4) || defined(MULMAT_QUANT)) && !defined(USE_OCP_FP4) -#ifdef MULMAT_QUANT -shared float ue4m3_fp32_lut[MmTypeA == GGML_TYPE_NVFP4 ? 128 : 8]; -#else +#if defined(DATA_A_NVFP4) && !defined(USE_OCP_FP4) shared float ue4m3_fp32_lut[128]; -#endif float ue4m3_to_fp32_build(uint u) { if (u == 0u || u == 127u) { @@ -1955,7 +1932,7 @@ float e8m0_to_fp32(uint8_t x) { return uintBitsToFloat(bits); } -#if defined(DATA_A_NVFP4) || defined(MULMAT_QUANT) +#if defined(DATA_A_NVFP4) #if defined(USE_OCP_FP4) floate4m3_t ue4m3_from_bits(uint8_t x) { if (x == uint8_t(0x7F)) { 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 8f120fd83d..10ced1e5af 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -245,6 +245,17 @@ bool is_iq_quant(const std::string& type_name) { return string_starts_with(type_name, "iq"); } +bool is_lut_quant(const std::string& type_name) { + return is_iq_quant(type_name) || type_name == "mxfp4" || type_name == "nvfp4"; +} + +std::string lut_load_vec_a(const std::string& type_name) { + if (type_name == "iq1_s" || type_name == "iq1_m" || type_name == "iq2_xxs" || type_name == "iq2_xs" || type_name == "iq2_s") { + return "8"; + } + return "4"; +} + static const char path_separator = '/'; std::string join_paths(const std::string& path1, const std::string& path2) { @@ -605,7 +616,6 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c continue; } - // Quant types: only MMQ (mul_mmq.comp) stays per-type std::string data_a_key = "DATA_A_" + to_uppercase(tname); const std::map float_type_dict = { {"FLOAT_TYPE", FLOAT_TYPE(1, tname)}, @@ -619,6 +629,26 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"},}), fp16, coopmat, coopmat2, f16acc); } #endif + + if (is_lut_quant(tname)) { + std::string lva = lut_load_vec_a(tname); + + string_to_spv(shader_name + "_" + tname + "_f16" + dot2_sfx, source_name, merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"LOAD_VEC_A", lva}, {"LOAD_VEC_B", load_vec}, {"B_TYPE", aligned_b_type_f16}, {"B_TYPE_SCALAR", "float16_t"}, {"B_TYPEV4", "f16vec4"}, {"D_TYPE", "float"}}), fp16, coopmat, coopmat2, f16acc); + + if (!coopmat2) { + string_to_spv(shader_name + "_" + tname + "_f32" + dot2_sfx, source_name, merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"LOAD_VEC_A", lva}, {"LOAD_VEC_B", load_vec}, {"B_TYPE", aligned_b_type_f32}, {"B_TYPE_SCALAR", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}), fp16, coopmat, coopmat2, f16acc); + } + +#if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) + if ((tname == "mxfp4" || tname == "nvfp4") && (coopmat || coopmat2)) { + string_to_spv(shader_name + "_" + tname + "_f16_ocp" + dot2_sfx, source_name, merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"USE_OCP_FP4", "1"}, {"LOAD_VEC_A", lva}, {"LOAD_VEC_B", load_vec}, {"B_TYPE", aligned_b_type_f16}, {"B_TYPE_SCALAR", "float16_t"}, {"B_TYPEV4", "f16vec4"}, {"D_TYPE", "float"}}), fp16, coopmat, coopmat2, f16acc); + if (!coopmat2) { + string_to_spv(shader_name + "_" + tname + "_f32_ocp" + dot2_sfx, source_name, merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"USE_OCP_FP4", "1"}, {"LOAD_VEC_A", lva}, {"LOAD_VEC_B", load_vec}, {"B_TYPE", aligned_b_type_f32}, {"B_TYPE_SCALAR", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}), fp16, coopmat, coopmat2, f16acc); + } + } +#endif + continue; + } } // Quant shader: one SPIR-V for all quant types, selected via MmTypeA spec constant