vulkan: add fp8 and scaled matmul support

This commit is contained in:
Ruben Ortlam
2026-09-14 13:56:07 +02:00
parent 8e99703d75
commit 3d30cf09f7
13 changed files with 311 additions and 8 deletions
+120 -4
View File
@@ -980,6 +980,7 @@ struct vk_device_struct {
vk_pipeline pipeline_concat_i8, pipeline_concat_i16, pipeline_concat_i32, pipeline_concat_i64;
vk_pipeline pipeline_upscale_nearest_f32, pipeline_upscale_bilinear_f32, pipeline_upscale_bicubic_f32, pipeline_upscale_bilinear_antialias_f32;
vk_pipeline pipeline_scale_f32;
vk_pipeline pipeline_mul_mat_scale_f32;
vk_pipeline pipeline_log[2];
vk_pipeline pipeline_tri[2];
vk_pipeline pipeline_diag[2];
@@ -1317,6 +1318,9 @@ struct vk_mat_mat_push_constants {
#define MAT_VEC_FUSION_FLAGS_BIAS1 0x2
#define MAT_VEC_FUSION_FLAGS_SCALE0 0x4
#define MAT_VEC_FUSION_FLAGS_SCALE1 0x8
#define MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE 0x10
#define MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_VEC 0x20
#define MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_2D 0x40
struct vk_mat_vec_push_constants {
uint32_t ncols;
@@ -1642,6 +1646,13 @@ struct vk_op_binary_push_constants {
float param1; float param2; int32_t param3;
};
struct vk_op_mul_mat_scale_push_constants {
uint32_t ne0; uint32_t ne1; uint32_t ne2; uint32_t ne3;
uint32_t sne0; uint32_t sne1; uint32_t sne2; uint32_t sne3;
uint32_t per_expert;
uint32_t ids_s0; uint32_t ids_s1;
};
// Distinct type with the same layout so concat can overload tensor offset initialization.
struct vk_op_concat_push_constants : vk_op_binary_push_constants {};
static_assert(sizeof(vk_op_concat_push_constants) == sizeof(vk_op_binary_push_constants));
@@ -4750,6 +4761,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_MXFP4], matmul_mxfp4_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3)
CREATE_MM2(pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_NVFP4], matmul_nvfp4_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3)
}
CREATE_MM2(pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_F8_E4M3], matmul_f8_e4m3_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3)
GGML_ASSERT(device->subgroup_ballot);
@@ -4791,6 +4803,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5)
CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5)
}
CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_F8_E4M3], matmul_id_subgroup_f8_e4m3_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5)
#undef CREATE_MM
#undef CREATE_MM2
} else
@@ -4865,6 +4878,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_MXFP4], matmul_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, );
CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, );
}
CREATE_MM2(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat[GGML_TYPE_F8_E4M3], matmul_f8_e4m3_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, );
GGML_ASSERT(device->subgroup_ballot);
@@ -4909,6 +4923,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
}
CREATE_MM2(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_F8_E4M3], matmul_id_subgroup_f8_e4m3_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id);
#undef CREATE_MM2
#undef CREATE_MM
} else
@@ -4992,6 +5007,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_NL], matmul_iq4_nl_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_MXFP4], matmul_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4], matmul_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
CREATE_MM2(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat[GGML_TYPE_F8_E4M3], matmul_f8_e4m3_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
if (device->integer_dot_product) {
@@ -5041,6 +5057,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
CREATE_MM2(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_F8_E4M3], matmul_id_subgroup_f8_e4m3_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
if (device->integer_dot_product) {
@@ -5089,6 +5106,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL], matmul_id_iq4_nl_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_mxfp4_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_nvfp4_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM2(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_F8_E4M3], matmul_id_f8_e4m3_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
if (device->integer_dot_product) {
@@ -5169,6 +5187,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_NL].f32acc, matmul_iq4_nl_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
CREATE_MM(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_MXFP4].f32acc, matmul_mxfp4_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
CREATE_MM(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat[GGML_TYPE_NVFP4].f32acc, matmul_nvfp4_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
CREATE_MM(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat[GGML_TYPE_F8_E4M3].f32acc, matmul_f8_e4m3_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
if (device->integer_dot_product) {
@@ -5217,6 +5236,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL].f32acc, matmul_id_subgroup_iq4_nl_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
CREATE_MM(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4].f32acc, matmul_id_subgroup_mxfp4_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
CREATE_MM(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4].f32acc, matmul_id_subgroup_nvfp4_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
CREATE_MM(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_F8_E4M3].f32acc, matmul_id_subgroup_f8_e4m3_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
} else {
CREATE_MM(GGML_TYPE_F32, pipeline_matmul_id_f32, matmul_id_f32_f32, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM(GGML_TYPE_F16, pipeline_matmul_id_f16.f32acc, matmul_id_f16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
@@ -5247,6 +5267,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL].f32acc, matmul_id_iq4_nl_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4].f32acc, matmul_id_mxfp4_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4].f32acc, matmul_id_nvfp4_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
CREATE_MM(GGML_TYPE_F8_E4M3, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_F8_E4M3].f32acc, matmul_id_f8_e4m3_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
}
}
// reusing CREATE_MM from the fp32 path
@@ -5356,6 +5377,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_IQ4_NL][i], "mul_mat_vec_iq4_nl_f32_f32", arr_dmmv_iq4_nl_f32_f32_len[reduc16], arr_dmmv_iq4_nl_f32_f32_data[reduc16], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_MXFP4][i], "mul_mat_vec_mxfp4_f32_f32", OCP_DMMV_LEN(arr_dmmv_mxfp4_f32_f32, reduc16), OCP_DMMV_DATA(arr_dmmv_mxfp4_f32_f32, reduc16), "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_NVFP4][i], "mul_mat_vec_nvfp4_f32_f32", OCP_DMMV_LEN(arr_dmmv_nvfp4_f32_f32, reduc16), OCP_DMMV_DATA(arr_dmmv_nvfp4_f32_f32, reduc16), "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_F8_E4M3][i], "mul_mat_vec_f8_e4m3_f32_f32", arr_dmmv_f8_e4m3_f32_f32_len[reduc16], arr_dmmv_f8_e4m3_f32_f32_data[reduc16], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_F32 ][i], "mul_mat_vec_f32_f16_f32", arr_dmmv_f32_f16_f32_len[reduc], arr_dmmv_f32_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1, 1, 1}, {wg_size_subgroup, 1, i+1}, 1, false, use_subgroups, force_subgroup_size);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_F16 ][i], "mul_mat_vec_f16_f16_f32", arr_dmmv_f16_f16_f32_len[reduc], arr_dmmv_f16_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2, 1, 1}, {wg_size_subgroup, 2, i+1}, 1, false, use_subgroups, force_subgroup_size);
@@ -5384,6 +5406,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_IQ4_NL][i], "mul_mat_vec_iq4_nl_f16_f32", arr_dmmv_iq4_nl_f16_f32_len[reduc16], arr_dmmv_iq4_nl_f16_f32_data[reduc16], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_MXFP4][i], "mul_mat_vec_mxfp4_f16_f32", OCP_DMMV_LEN(arr_dmmv_mxfp4_f16_f32, reduc16), OCP_DMMV_DATA(arr_dmmv_mxfp4_f16_f32, reduc16), "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_NVFP4][i], "mul_mat_vec_nvfp4_f16_f32", OCP_DMMV_LEN(arr_dmmv_nvfp4_f16_f32, reduc16), OCP_DMMV_DATA(arr_dmmv_nvfp4_f16_f32, reduc16), "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_F8_E4M3][i], "mul_mat_vec_f8_e4m3_f16_f32", arr_dmmv_f8_e4m3_f16_f32_len[reduc16], arr_dmmv_f8_e4m3_f16_f32_data[reduc16], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq, i+1}, 1, true, use_subgroups16, force_subgroup_size16);
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
if (device->integer_dot_product) {
@@ -5439,6 +5462,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_IQ4_NL], "mul_mat_vec_id_iq4_nl_f32", arr_dmmv_id_iq4_nl_f32_f32_len[reduc16], arr_dmmv_id_iq4_nl_f32_f32_data[reduc16], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_MXFP4], "mul_mat_vec_id_mxfp4_f32", OCP_DMMV_LEN(arr_dmmv_id_mxfp4_f32_f32, reduc16), OCP_DMMV_DATA(arr_dmmv_id_mxfp4_f32_f32, reduc16), "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_NVFP4], "mul_mat_vec_id_nvfp4_f32", OCP_DMMV_LEN(arr_dmmv_id_nvfp4_f32_f32, reduc16), OCP_DMMV_DATA(arr_dmmv_id_nvfp4_f32_f32, reduc16), "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq}, 1, true, use_subgroups16, force_subgroup_size16);
ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_F8_E4M3], "mul_mat_vec_id_f8_e4m3_f32", arr_dmmv_id_f8_e4m3_f32_f32_len[reduc16], arr_dmmv_id_f8_e4m3_f32_f32_data[reduc16], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {wg_size_subgroup16, rm_iq}, 1, true, use_subgroups16, force_subgroup_size16);
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
if (device->integer_dot_product) {
@@ -5505,6 +5529,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_IQ4_NL], "dequant_iq4_nl", dequant_iq4_nl_len, dequant_iq4_nl_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_MXFP4], "dequant_mxfp4", dequant_mxfp4_len, dequant_mxfp4_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_NVFP4], "dequant_nvfp4", dequant_nvfp4_len, dequant_nvfp4_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_F8_E4M3], "dequant_f8_e4m3", dequant_f8_e4m3_len, dequant_f8_e4m3_data, "main", 2, 5 * sizeof(uint32_t), {256, 1, 1}, {}, 1);
// get_rows
ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_F32 ], "get_rows_f32", get_rows_f32_len, get_rows_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), { 512, 1, 1}, {}, 1);
@@ -5534,6 +5559,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_IQ4_NL], "get_rows_iq4_nl", get_rows_iq4_nl_len, get_rows_iq4_nl_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_MXFP4], "get_rows_mxfp4", get_rows_mxfp4_len, get_rows_mxfp4_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_NVFP4], "get_rows_nvfp4", get_rows_nvfp4_len, get_rows_nvfp4_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_F8_E4M3], "get_rows_f8_e4m3", get_rows_f8_e4m3_len, get_rows_f8_e4m3_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_I32], "get_rows_i32", get_rows_i32_len, get_rows_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_F32 ], "get_rows_f32_f32", get_rows_f32_f32_len, get_rows_f32_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), { 512, 1, 1}, {}, 1);
@@ -5563,6 +5589,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_IQ4_NL], "get_rows_iq4_nl_f32", get_rows_iq4_nl_f32_len, get_rows_iq4_nl_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_MXFP4], "get_rows_mxfp4_f32", get_rows_mxfp4_f32_len, get_rows_mxfp4_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_NVFP4], "get_rows_nvfp4_f32", get_rows_nvfp4_f32_len, get_rows_nvfp4_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_F8_E4M3], "get_rows_f8_e4m3_f32", get_rows_f8_e4m3_f32_len, get_rows_f8_e4m3_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_get_rows_back_f32, "get_rows_back_f32", get_rows_back_f32_len, get_rows_back_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {256, 1, 1}, {}, 1, true);
ggml_vk_create_pipeline(device, device->pipeline_matmul_split_k_reduce, "split_k_reduce", split_k_reduce_len, split_k_reduce_data, "main", 2, 2 * sizeof(uint32_t), {256 * 4, 1, 1}, {}, 1);
@@ -5715,6 +5742,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_scale_f32, "scale_f32", scale_f32_len, scale_f32_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_mul_mat_scale_f32, "mul_mat_scale_f32", mul_mat_scale_f32_len, mul_mat_scale_f32_data, "main", 3, sizeof(vk_op_mul_mat_scale_push_constants), {256, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_log[0], "log_f32", log_f32_len, log_f32_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_log[1], "log_f16", log_f16_len, log_f16_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1);
@@ -7787,6 +7816,7 @@ static vk_pipeline ggml_vk_get_to_fp16(ggml_backend_vk_context * ctx, ggml_type
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
break;
default:
return nullptr;
@@ -7862,6 +7892,7 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
break;
default:
return nullptr;
@@ -7932,6 +7963,7 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec(ggml_backend_vk_context *
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
break;
default:
return nullptr;
@@ -8026,6 +8058,7 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
break;
default:
return nullptr;
@@ -8099,6 +8132,7 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec_id(ggml_backend_vk_context
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
break;
default:
return nullptr;
@@ -9839,6 +9873,15 @@ static void ggml_vk_mul_mat_vec_q_f16(ggml_backend_vk_context * ctx, vk_context&
fusion_flags |= MAT_VEC_FUSION_FLAGS_BIAS1;
}
// fold weight scale (src[2]) via the F0 slot, free since scaled matmuls don't fuse
if (dst->src[2] != nullptr) {
d_F0 = ggml_vk_tensor_subbuffer(ctx, dst->src[2]);
fusion_flags |= MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE;
if (ggml_nelements(dst->src[2]) > 1) {
fusion_flags |= MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_VEC;
}
}
ggml_pipeline_request_descriptor_sets(ctx, dmmv, CEIL_DIV(ne12 * ne13, ctx->device->properties.limits.maxComputeWorkGroupCount[1]));
uint32_t base_work_group_y = 0;
@@ -10119,12 +10162,49 @@ static void ggml_vk_fwht(ggml_backend_vk_context * ctx, vk_context& subctx, cons
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src_buf, dst_buf }, pc, { workgroups_x, 1, 1 });
}
static void ggml_vk_mul_mat_scale(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) {
const bool is_id = dst->op == GGML_OP_MUL_MAT_ID;
const ggml_tensor * scale = is_id ? dst->src[3] : dst->src[2];
const ggml_tensor * ids = is_id ? dst->src[2] : nullptr;
GGML_ASSERT(scale != nullptr && scale->type == GGML_TYPE_F32 && ggml_is_contiguous(scale));
GGML_ASSERT(dst->type == GGML_TYPE_F32 && ggml_is_contiguous(dst));
const bool per_expert = is_id && ((uint64_t)ggml_nelements(scale) == (uint64_t)dst->src[0]->ne[2] ||
(uint64_t)scale->ne[1] == (uint64_t)dst->src[0]->ne[2]);
vk_pipeline pipeline = ctx->device->pipeline_mul_mat_scale_f32;
ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);
ggml_vk_sync_buffers(ctx, subctx);
vk_subbuffer d_D = ggml_vk_tensor_subbuffer(ctx, dst);
vk_subbuffer d_S = ggml_vk_tensor_subbuffer(ctx, scale);
vk_subbuffer d_I = per_expert ? ggml_vk_tensor_subbuffer(ctx, ids) : d_S;
const vk_op_mul_mat_scale_push_constants pc = {
(uint32_t)dst->ne[0], (uint32_t)dst->ne[1], (uint32_t)dst->ne[2], (uint32_t)dst->ne[3],
(uint32_t)scale->ne[0], (uint32_t)scale->ne[1], (uint32_t)scale->ne[2], (uint32_t)scale->ne[3],
per_expert ? 1u : 0u,
per_expert ? (uint32_t)(ids->nb[0] / sizeof(int32_t)) : 0u,
per_expert ? (uint32_t)(ids->nb[1] / sizeof(int32_t)) : 0u,
};
const uint64_t total = (uint64_t)ggml_nelements(dst);
const uint64_t max_elems = (uint64_t)ctx->device->properties.limits.maxComputeWorkGroupCount[0] * 256;
const uint32_t elements0 = (uint32_t)std::min(total, max_elems);
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { d_D, d_S, d_I }, pc, { elements0, 1, 1 });
}
static void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx) {
ggml_tensor * dst = cgraph->nodes[node_idx];
ggml_tensor * src0 = dst->src[0];
ggml_tensor * src1 = dst->src[1];
VK_LOG_DEBUG("ggml_vk_mul_mat(" << src0 << ", " << src1 << ", " << dst << ")");
bool folded = false;
// Handle huge A matrix by splitting the M dimensions. This works well for convolution use cases
// where the M dimension is very large.
// Split_k doesn't work with M splitting.
@@ -10176,11 +10256,16 @@ static void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, c
// mul_mat_vec supports batching ne12*ne13 when ne11==1, or treating ne11 as the batch size (up to four)
// when ne12 and ne13 are one.
} else if ((dst->ne[1] == 1 || (dst->ne[1] <= mul_mat_vec_max_cols && src1->ne[2] * src1->ne[3] == 1)) &&
(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16 || ggml_is_quantized(src0->type))) {
(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16 || src0->type == GGML_TYPE_F8_E4M3 || ggml_is_quantized(src0->type))) {
ggml_vk_mul_mat_vec_q_f16(ctx, subctx, cgraph, node_idx);
folded = true;
} else {
ggml_vk_mul_mat_q_f16(ctx, subctx, src0, src1, dst, false);
}
if (dst->src[2] != nullptr && !folded) {
ggml_vk_mul_mat_scale(ctx, subctx, dst);
}
}
static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst) {
@@ -10755,6 +10840,18 @@ static void ggml_vk_mul_mat_vec_id_q_f16(ggml_backend_vk_context * ctx, vk_conte
fusion_flags |= MAT_VEC_FUSION_FLAGS_SCALE1;
}
// fold weight scale (src[3]) via the F0 slot, free since scaled matmuls don't fuse
if (dst->src[3] != nullptr) {
d_F0 = ggml_vk_tensor_subbuffer(ctx, dst->src[3]);
fusion_flags |= MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE;
if (ggml_nelements(dst->src[3]) > 1) {
fusion_flags |= MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_VEC;
}
if (dst->src[3]->ne[1] > 1) {
fusion_flags |= MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_2D;
}
}
// Loop over the batch dimension
for (uint32_t expert_i1 = 0; expert_i1 < nei1; ++expert_i1) {
const vk_mat_vec_id_push_constants pc = {
@@ -10787,7 +10884,7 @@ static bool ggml_vk_use_mul_mat_vec_id(const struct ggml_cgraph * cgraph, int no
ggml_tensor * dst = cgraph->nodes[node_idx];
ggml_tensor * src0 = dst->src[0];
ggml_tensor * src2 = dst->src[2];
return (src2->ne[1] <= 8) && (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type));
return (src2->ne[1] <= 8) && (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_F8_E4M3 || ggml_is_quantized(src0->type));
}
static void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx) {
@@ -10796,11 +10893,17 @@ static void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx
ggml_tensor * src1 = dst->src[1];
ggml_tensor * src2 = dst->src[2];
VK_LOG_DEBUG("ggml_vk_mul_mat_id(" << src0 << ", " << src1 << ", " << src2 << ", " << dst << ")");
bool folded = false;
if (ggml_vk_use_mul_mat_vec_id(cgraph, node_idx)) {
ggml_vk_mul_mat_vec_id_q_f16(ctx, subctx, cgraph, node_idx);
folded = true;
} else {
ggml_vk_mul_mat_id_q_f16(ctx, subctx, src0, src1, src2, dst);
}
if (dst->src[3] != nullptr && !folded) {
ggml_vk_mul_mat_scale(ctx, subctx, dst);
}
}
static bool ggml_vk_flash_attn_scalar_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type, ggml_type v_type) {
@@ -16917,6 +17020,15 @@ static bool ggml_vk_can_fuse(const ggml_backend_vk_context * ctx, const struct g
return false;
}
// don't fuse onto a scaled matmul: an additive bias would apply before the scale epilogue
{
const ggml_tensor * base = cgraph->nodes[node_idx];
if ((base->op == GGML_OP_MUL_MAT && base->src[2] != nullptr) ||
(base->op == GGML_OP_MUL_MAT_ID && base->src[3] != nullptr)) {
return false;
}
}
if (ops.size() == 2 && ops.begin()[0] == GGML_OP_RMS_NORM && ops.begin()[1] == GGML_OP_MUL) {
// additional constraints specific to this fusion
const ggml_tensor *rms_norm = cgraph->nodes[node_idx];
@@ -18673,6 +18785,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
break;
default:
return false;
@@ -18779,6 +18892,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_TQ2_0:
case GGML_TYPE_F8_E4M3:
case GGML_TYPE_I32:
return true;
default:
@@ -19762,9 +19876,11 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph *
ggml_flash_attn_ext_add_sinks(tensor_clone, src_clone[4]);
}
} else if (tensor->op == GGML_OP_MUL_MAT) {
tensor_clone = ggml_mul_mat(ggml_ctx, src_clone[0], src_clone[1]);
tensor_clone = src_clone[2] ? ggml_mul_mat_ext(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], nullptr)
: ggml_mul_mat(ggml_ctx, src_clone[0], src_clone[1]);
} else if (tensor->op == GGML_OP_MUL_MAT_ID) {
tensor_clone = ggml_mul_mat_id(ggml_ctx, src_clone[0], src_clone[1], src_clone[2]);
tensor_clone = src_clone[3] ? ggml_mul_mat_id_ext(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3], nullptr)
: ggml_mul_mat_id(ggml_ctx, src_clone[0], src_clone[1], src_clone[2]);
} else if (tensor->op == GGML_OP_SUB) {
tensor_clone = ggml_sub(ggml_ctx, src_clone[0], src_clone[1]);
} else if (tensor->op == GGML_OP_MUL) {
@@ -0,0 +1,18 @@
#version 450
#include "dequant_head.glsl"
layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in;
layout (binding = 0) readonly buffer A {uint8_t data_a[];};
layout (binding = 1) writeonly buffer D {D_TYPE data_b[];};
void main() {
const uint i = gl_GlobalInvocationID.x;
if (i >= p.nel) {
return;
}
data_b[i] = D_TYPE(e4m3_to_fp32(data_a[i]));
}
@@ -543,12 +543,38 @@ vec4 dequantize4(uint ib, uint iqs, uint a_offset) {
}
#endif
#if defined(DATA_A_F8_E4M3)
FLOAT_TYPE dequantize1(uint ib, uint iqs, uint a_offset) {
return FLOAT_TYPE(e4m3_to_fp32(data_a[a_offset + ib + iqs]));
}
vec2 dequantize(uint ib, uint iqs, uint a_offset) {
const uint idx = a_offset + ib + iqs;
return vec2(e4m3_to_fp32(data_a[idx]), e4m3_to_fp32(data_a[idx + 1]));
}
vec4 dequantize4(uint ib, uint iqs, uint a_offset) {
const uint idx = a_offset + ib + iqs;
return vec4(e4m3_to_fp32(data_a[idx ]), e4m3_to_fp32(data_a[idx + 1]),
e4m3_to_fp32(data_a[idx + 2]), e4m3_to_fp32(data_a[idx + 3]));
}
vec4 dequantize4_2aligned(uint ib, uint iqs, uint a_offset) {
const uint idx = a_offset + ib + iqs;
return vec4(e4m3_to_fp32(data_a[idx ]), e4m3_to_fp32(data_a[idx + 1]),
e4m3_to_fp32(data_a[idx + 2]), e4m3_to_fp32(data_a[idx + 3]));
}
#endif
#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)
vec2 get_dm(uint ib, uint a_offset) {
return vec2(0, 0);
}
#endif
#if defined(DATA_A_F8_E4M3)
vec2 get_dm(uint ib, uint a_offset) {
return vec2(1.0, 0.0);
}
#endif
#if defined(DATA_A_IQ1_M)
vec2 get_dm(uint ib, uint a_offset) {
const uint16_t[4] scales = data_a[a_offset + ib].scales;
@@ -1385,6 +1385,17 @@ f16vec4 dequantFuncNVFP4_v(const in decodeBufNVFP4 bl, const in uint blockCoords
}
#endif
#if defined(DATA_A_F8_E4M3)
layout(buffer_reference, std430, buffer_reference_align = 1) buffer decodeBufF8_E4M3 {
uint8_t block;
};
float16_t dequantFuncF8_E4M3(const in decodeBufF8_E4M3 bl, const in uint blockCoords[2], const in uint coordInBlock[2])
{
return float16_t(e4m3_to_fp32(bl.block));
}
#endif
#if defined(DATA_A_Q1_0)
#define dequantFuncA dequantFuncQ1_0
#define dequantFuncA_v dequantFuncQ1_0_v
@@ -1463,4 +1474,6 @@ f16vec4 dequantFuncNVFP4_v(const in decodeBufNVFP4 bl, const in uint blockCoords
#define dequantFuncA_v dequantFuncNVFP4_v
#elif defined(DATA_A_F32)
#define dequantFuncA dequantFuncF32
#elif defined(DATA_A_F8_E4M3)
#define dequantFuncA dequantFuncF8_E4M3
#endif
@@ -0,0 +1,41 @@
#version 450
#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require
layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in;
layout(binding = 0) buffer D { float data_d[]; };
layout(binding = 1) readonly buffer S { float data_s[]; };
layout(binding = 2) readonly buffer I { int32_t data_ids[]; };
layout(push_constant) uniform parameter {
uint ne0; uint ne1; uint ne2; uint ne3;
uint sne0; uint sne1; uint sne2; uint sne3;
uint per_expert;
uint ids_s0; uint ids_s1;
} p;
void main() {
const uint total = p.ne0 * p.ne1 * p.ne2 * p.ne3;
const uint stride = gl_NumWorkGroups.x * gl_WorkGroupSize.x;
for (uint i = gl_GlobalInvocationID.x; i < total; i += stride) {
const uint i0 = i % p.ne0;
const uint i1 = (i / p.ne0) % p.ne1;
const uint i2 = (i / (p.ne0 * p.ne1)) % p.ne2;
const uint i3 = i / (p.ne0 * p.ne1 * p.ne2);
uint si;
if (p.per_expert != 0) {
const uint expert = uint(data_ids[i1 * p.ids_s0 + i2 * p.ids_s1]);
si = (p.sne1 > 1) ? (i0 % p.sne0) + expert * p.sne0 : expert;
} else {
si = (i0 % p.sne0)
+ (i1 % p.sne1) * p.sne0
+ (i2 % p.sne2) * p.sne0 * p.sne1
+ (i3 % p.sne3) * p.sne0 * p.sne1 * p.sne2;
}
data_d[i] *= data_s[si];
}
}
@@ -7,7 +7,7 @@
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
#if !defined(DATA_A_F32) && !defined(DATA_A_F16) && !defined(DATA_A_BF16)
#if !defined(DATA_A_F32) && !defined(DATA_A_F16) && !defined(DATA_A_BF16) && !defined(DATA_A_F8_E4M3)
#define K_PER_ITER 8
#else
#define K_PER_ITER 4
@@ -114,7 +114,7 @@ void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const
}
}
#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)
#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16) || defined(DATA_A_F8_E4M3)
void iter_aligned_nonquant(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint tid, const uint i)
{
[[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
@@ -90,6 +90,16 @@ layout (constant_id = 0) const uint BLOCK_SIZE = 32;
layout (constant_id = 1) const uint NUM_ROWS = 1;
layout (constant_id = 2) const uint NUM_COLS = 1;
#ifdef MUL_MAT_ID
#define WEIGHT_SCALE_INDEX(row) (((p.fusion_flags & MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_2D) != 0) ? (expert_id * p.stride_d + (row)) : (((p.fusion_flags & MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_VEC) != 0) ? expert_id : 0u))
#else
#define WEIGHT_SCALE_INDEX(row) (((p.fusion_flags & MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_VEC) != 0) ? (row) : 0u)
#endif
#define APPLY_WEIGHT_SCALE(acc, row) \
if ((p.fusion_flags & MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE) != 0) { \
(acc) *= FLOAT_TYPE(data_fuse0[WEIGHT_SCALE_INDEX(row)]); \
}
#ifdef USE_SUBGROUP_ADD_NO_SHMEM
void reduce_result(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const in uint32_t d_offset, const in uint32_t first_row, const in uint32_t num_rows, const in uint32_t tid) {
[[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
@@ -101,6 +111,7 @@ void reduce_result(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const in uint32_t
if (tid == 0) {
[[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
[[unroll]] for (uint n = 0; n < num_rows; ++n) {
APPLY_WEIGHT_SCALE(temp[j][n], first_row + n);
#ifdef MUL_MAT_ID
if ((p.fusion_flags & MAT_VEC_FUSION_FLAGS_BIAS0) != 0) {
temp[j][n] += FLOAT_TYPE(data_fuse0[expert_id*p.stride_d + first_row + n]);
@@ -156,6 +167,7 @@ void reduce_result(FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const in uint32_t d_offs
[[unroll]] for (uint s = 0; s < gl_NumSubgroups; ++s) {
temp[j][n] += tmpsh[j][n][s];
}
APPLY_WEIGHT_SCALE(temp[j][n], first_row + n);
#ifdef MUL_MAT_ID
if ((p.fusion_flags & MAT_VEC_FUSION_FLAGS_BIAS0) != 0) {
temp[j][n] += FLOAT_TYPE(data_fuse0[expert_id*p.stride_d + first_row + n]);
@@ -201,6 +213,7 @@ void reduce_result(FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const in uint32_t d_offs
if (tid == 0) {
[[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
[[unroll]] for (uint n = 0; n < num_rows; ++n) {
APPLY_WEIGHT_SCALE(tmpsh[j][n][0], first_row + n);
#ifdef MUL_MAT_ID
if ((p.fusion_flags & MAT_VEC_FUSION_FLAGS_BIAS0) != 0) {
tmpsh[j][n][0] += FLOAT_TYPE(data_fuse0[expert_id*p.stride_d + first_row + n]);
@@ -4,6 +4,11 @@
#define MAT_VEC_FUSION_FLAGS_BIAS1 0x2
#define MAT_VEC_FUSION_FLAGS_SCALE0 0x4
#define MAT_VEC_FUSION_FLAGS_SCALE1 0x8
// weight scale in data_fuse0; _VEC: per-output (dense) / per-expert (id), else scalar;
// _2D (id only): per-channel-per-expert [n_out, n_expert]
#define MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE 0x10
#define MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_VEC 0x20
#define MAT_VEC_FUSION_FLAGS_WEIGHT_SCALE_2D 0x40
layout (binding = 0) readonly buffer A {A_TYPE data_a[];};
#if defined(A_TYPEV4)
@@ -53,6 +53,8 @@ layout (binding = 0) readonly buffer A_SCALAR {float data_a_scalar[];};
layout (binding = 0) readonly buffer A_SCALAR {float16_t data_a_scalar[];};
#elif defined(DATA_A_BF16)
layout (binding = 0) readonly buffer A_SCALAR {uint16_t data_a_scalar[];};
#elif defined(DATA_A_F8_E4M3)
layout (binding = 0) readonly buffer A_SCALAR {uint8_t data_a_scalar[];};
#endif
#if defined(A_TYPE_PACKED16)
layout (binding = 0) readonly buffer A_PACKED16 {A_TYPE_PACKED16 data_a_packed16[];};
@@ -197,7 +199,7 @@ void main() {
const uint warp_r = warp_i % (BM / WM);
const uint warp_c = warp_i / (BM / WM);
#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)
#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16) || defined(DATA_A_F8_E4M3)
const uint LOAD_VEC_A_EFF = (ALIGNED != 0) ? LOAD_VEC_A : 1;
const uint LOAD_VEC_BATCH_A = (ALIGNED != 0) ? 1 : 2;
#else
@@ -79,7 +79,7 @@ layout (binding = 2) writeonly buffer D {D_TYPE data_d[];};
layout (binding = 1) readonly buffer B4 {B_TYPEV4 data_b_v4[];};
#endif
#if QUANT_K > 1
#if QUANT_K > 1 || defined(DATA_A_F8_E4M3)
#include "dequant_funcs_cm2.glsl"
#if defined(dequantFuncA_v) && defined(GGML_VULKAN_COOPMAT2_DECODE_VECTOR)
#define DECODEFUNCA , dequantFuncA, dequantFuncA_v
@@ -580,6 +580,26 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin
store_a(col, eff_row + 4, FLOAT_TYPEV2(kvalues_mxfp4[vui >> 4] * d,
kvalues_mxfp4[vui2 >> 4] * d));
#endif
#elif defined(DATA_A_F8_E4M3)
#if LOAD_VEC_A == 4
if (ALIGNED != 0) {
const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
const uint k_pair = row * LOAD_VEC_A / 2;
FLOAT_TYPEV4 aa = FLOAT_TYPEV4(e4m3_to_fp32(data_a[idx]));
store_a(col, k_pair, aa.xy);
store_a(col, k_pair + 1, aa.zw);
return;
}
#endif
const uint idx = pos_a + col * p.stride_a + row * 2;
if (idx_m < p.M && block + row * 2 + 1 < end_k) {
store_a(col, row, FLOAT_TYPEV2(e4m3_to_fp32(data_a_scalar[idx]),
e4m3_to_fp32(data_a_scalar[idx + 1])));
} else if (idx_m < p.M && block + row * 2 < end_k) {
store_a(col, row, FLOAT_TYPEV2(e4m3_to_fp32(data_a_scalar[idx]), 0.0f));
} else {
store_a(col, row, FLOAT_TYPEV2(0.0f));
}
#endif
}
@@ -12,6 +12,10 @@
#extension GL_EXT_float_e4m3 : require
#endif
#if defined(DATA_A_F8_E4M3)
#extension GL_EXT_shader_8bit_storage : require
#endif
#if defined(DATA_A_F32)
#define QUANT_K 1
#define QUANT_R 1
@@ -53,6 +57,18 @@
#define A_TYPE_PACKED32 uint32_t
#endif
#if defined(DATA_A_F8_E4M3)
#define QUANT_K 1
#define QUANT_R 1
#if LOAD_VEC_A == 4
#define A_TYPE u8vec4
#else
#define A_TYPE uint8_t
#endif
#define A_TYPE_PACKED32 uint32_t
#endif
#define QUANT_K_Q4_0 32
#define QUANT_R_Q4_0 2
@@ -1883,6 +1899,31 @@ float bf16_to_fp32(uint32_t u)
return uintBitsToFloat(u << 16);
}
// OCP E4M3: 1 sign, 4 exponent (bias 7), 3 mantissa. No infinities; S.1111.111 is NaN.
float e4m3_to_fp32(uint x) {
const uint s = (x >> 7) & 1u;
const uint e = (x >> 3) & 0xFu;
const uint m = x & 0x7u;
float val;
if (e == 0u) {
val = float(m) * (1.0 / 512.0); // 2^-6 * m/8
} else {
val = uintBitsToFloat(((e + 120u) << 23) | (m << 20));
}
return s == 1u ? -val : val;
}
#if defined(DATA_A_F8_E4M3)
float e4m3_to_fp32(uint8_t x) {
return e4m3_to_fp32(uint(x));
}
vec4 e4m3_to_fp32(u8vec4 x) {
return vec4(e4m3_to_fp32(uint(x.x)), e4m3_to_fp32(uint(x.y)),
e4m3_to_fp32(uint(x.z)), e4m3_to_fp32(uint(x.w)));
}
#endif
vec4 bf16_to_fp32(uvec4 u)
{
return vec4(bf16_to_fp32(u.x), bf16_to_fp32(u.y), bf16_to_fp32(u.z), bf16_to_fp32(u.w));
@@ -74,6 +74,7 @@ const std::vector<std::string> type_names = {
"nvfp4",
"tq2_0",
"bf16",
"f8_e4m3",
};
enum MatMulIdType {
@@ -595,6 +596,11 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c
std::string data_a_key = "DATA_A_" + to_uppercase(tname);
// For aligned matmul loads
std::string load_vec_a = (coopmat2 || tname == "f32" || tname == "f16" || tname == "bf16") ? load_vec : load_vec_quant;
// f8_e4m3 is a 1-byte scalar float: load 4 bytes at a time (u8vec4) for the aligned path,
// except coopmat2 which decodes per-element via dequantFunc (needs uint8_t A_TYPE, LOAD_VEC_A=1).
if (tname == "f8_e4m3") {
load_vec_a = coopmat2 ? "1" : "4";
}
const std::map<std::string, std::string> float_type_dict = {
{"FLOAT_TYPE", FLOAT_TYPE(1, tname)},
@@ -895,6 +901,8 @@ void process_shaders() {
string_to_spv("scale_f32", "scale.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}});
string_to_spv("mul_mat_scale_f32", "mul_mat_scale.comp", {});
string_to_spv("pad_f32", "pad.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}});
string_to_spv("pad_reflect_1d_f32", "pad_reflect_1d.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}});