mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-08-31 17:17:44 +02:00
remove LUT quants from unified shader
This commit is contained in:
@@ -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<std::vector<uint32_t>(const std::vector<uint32_t>&, 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<uint32_t>& 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<uint32_t>& 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<vk_tile_config> 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<vk_tile_config> 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)
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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[];};
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)) {
|
||||
|
||||
@@ -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<std::string, std::string> 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
|
||||
|
||||
Reference in New Issue
Block a user