diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index a04a6b27a8..0d450cb7b8 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f8_e4m3.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f8_e4m3.comp new file mode 100644 index 0000000000..3f04e7c213 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f8_e4m3.comp @@ -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])); +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl index 627932bd35..b40b68758a 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl @@ -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; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl index 46cc69cb26..efec185c5f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl @@ -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 diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_scale.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_scale.comp new file mode 100644 index 0000000000..004a7d0abb --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_scale.comp @@ -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]; + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp index 5a9d0e778f..1f5044432e 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp @@ -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) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_base.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_base.glsl index 4aeda68c7f..e5688d3640 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_base.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_base.glsl @@ -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]); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl index e8d053cdd4..5ca24b4197 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl @@ -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) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp index 63c4aaebcb..f1c55edc1f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp @@ -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 diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp index 27f3178e7f..ac055bc06b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp @@ -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 diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl index 7d852dced8..e7e0b24ee8 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl @@ -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 } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl index adb1bb8b32..3e2df0577c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl @@ -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)); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 27ff68c10d..39e6421464 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -74,6 +74,7 @@ const std::vector 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 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"}});