Skip to content

Commit b4e3dc6

Browse files
authored
vulkan: add v_dot2_f32_f16 support in matrix-matrix multiplication and Flash Attention (ggml-org#24123)
* vulkan: add support for valve fp16 dot2 extension * use macro for dot2 path choice * properly check for the feature * add dot_product abstraction to reduce preprocessor branching
1 parent ae735b1 commit b4e3dc6

5 files changed

Lines changed: 139 additions & 35 deletions

File tree

ggml/src/ggml-vulkan/ggml-vulkan.cpp

Lines changed: 74 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -113,6 +113,21 @@ typedef struct VkPhysicalDeviceShaderBfloat16FeaturesKHR {
113113
} VkPhysicalDeviceShaderBfloat16FeaturesKHR;
114114
#endif
115115

116+
#if !defined(VK_VALVE_shader_mixed_float_dot_product)
117+
#define VK_VALVE_shader_mixed_float_dot_product 1
118+
#define VK_VALVE_SHADER_MIXED_FLOAT_DOT_PRODUCT_SPEC_VERSION 1
119+
#define VK_VALVE_SHADER_MIXED_FLOAT_DOT_PRODUCT_EXTENSION_NAME "VK_VALVE_shader_mixed_float_dot_product"
120+
#define VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_MIXED_FLOAT_DOT_PRODUCT_FEATURES_VALVE ((VkStructureType)1000673000)
121+
typedef struct VkPhysicalDeviceShaderMixedFloatDotProductFeaturesVALVE {
122+
VkStructureType sType;
123+
void* pNext;
124+
VkBool32 shaderMixedFloatDotProductFloat16AccFloat32;
125+
VkBool32 shaderMixedFloatDotProductFloat16AccFloat16;
126+
VkBool32 shaderMixedFloatDotProductBFloat16Acc;
127+
VkBool32 shaderMixedFloatDotProductFloat8AccFloat32;
128+
} VkPhysicalDeviceShaderMixedFloatDotProductFeaturesVALVE;
129+
#endif
130+
116131
#define ROUNDUP_POW2(M, N) (((M) + (N) - 1) & ~((N) - 1))
117132
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
118133
static bool is_pow2(uint32_t x) { return x > 1 && (x & (x-1)) == 0; }
@@ -705,6 +720,8 @@ struct vk_device_struct {
705720
bool coopmat2_bf16_support {};
706721
bool coopmat2_decode_vector;
707722

723+
bool dot2_f16 {};
724+
708725
bool pipeline_executable_properties_support {};
709726

710727
size_t idx;
@@ -3920,8 +3937,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
39203937
name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16";
39213938
} else {
39223939
if (device->fp16) {
3923-
if (f32acc) { spv_data = flash_attn_f32_f16_data; spv_size = flash_attn_f32_f16_len; }
3924-
else { spv_data = flash_attn_f32_f16_f16acc_data; spv_size = flash_attn_f32_f16_f16acc_len; }
3940+
if (device->dot2_f16) {
3941+
if (f32acc) { spv_data = flash_attn_f32_f16_dot2_data; spv_size = flash_attn_f32_f16_dot2_len; }
3942+
else { spv_data = flash_attn_f32_f16_dot2_f16acc_data; spv_size = flash_attn_f32_f16_dot2_f16acc_len; }
3943+
} else {
3944+
if (f32acc) { spv_data = flash_attn_f32_f16_data; spv_size = flash_attn_f32_f16_len; }
3945+
else { spv_data = flash_attn_f32_f16_f16acc_data; spv_size = flash_attn_f32_f16_f16acc_len; }
3946+
}
39253947
} else {
39263948
spv_data = flash_attn_f32_f16_fp32_data;
39273949
spv_size = flash_attn_f32_f16_fp32_len;
@@ -4215,7 +4237,23 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
42154237
#endif // defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT)
42164238
if (device->fp16) {
42174239
// Create 6 variants, {s,m,l}x{unaligned,aligned}
4240+
// Selects dot2 SPIR-V variant at runtime when device->dot2_f16 is true
42184241
#define CREATE_MM(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID, REQSUBGROUPSIZE) \
4242+
if (device->mul_mat ## ID ## _l[TYPE]) \
4243+
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", (device->dot2_f16 ? NAMELC ## _dot2 ## F16ACC ## _len : NAMELC ## F16ACC ## _len), (device->dot2_f16 ? NAMELC ## _dot2 ## F16ACC ## _data : NAMELC ## F16ACC ## _data), "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
4244+
if (device->mul_mat ## ID ## _m[TYPE]) \
4245+
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", (device->dot2_f16 ? NAMELC ## _dot2 ## F16ACC ## _len : NAMELC ## F16ACC ## _len), (device->dot2_f16 ? NAMELC ## _dot2 ## F16ACC ## _data : NAMELC ## F16ACC ## _data), "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
4246+
if (device->mul_mat ## ID ## _s[TYPE]) \
4247+
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", (device->dot2_f16 ? NAMELC ## _dot2 ## F16ACC ## _len : NAMELC ## F16ACC ## _len), (device->dot2_f16 ? NAMELC ## _dot2 ## F16ACC ## _data : NAMELC ## F16ACC ## _data), "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
4248+
if (device->mul_mat ## ID ## _l[TYPE]) \
4249+
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", (device->dot2_f16 ? NAMELC ## _dot2_aligned ## F16ACC ## _len : NAMELC ## _aligned ## F16ACC ## _len), (device->dot2_f16 ? NAMELC ## _dot2_aligned ## F16ACC ## _data : NAMELC ## _aligned ## F16ACC ## _data), "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, l_align, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
4250+
if (device->mul_mat ## ID ## _m[TYPE]) \
4251+
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", (device->dot2_f16 ? NAMELC ## _dot2_aligned ## F16ACC ## _len : NAMELC ## _aligned ## F16ACC ## _len), (device->dot2_f16 ? NAMELC ## _dot2_aligned ## F16ACC ## _data : NAMELC ## _aligned ## F16ACC ## _data), "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, m_align, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
4252+
if (device->mul_mat ## ID ## _s[TYPE]) \
4253+
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", (device->dot2_f16 ? NAMELC ## _dot2_aligned ## F16ACC ## _len : NAMELC ## _aligned ## F16ACC ## _len), (device->dot2_f16 ? NAMELC ## _dot2_aligned ## F16ACC ## _data : NAMELC ## _aligned ## F16ACC ## _data), "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, s_align, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
4254+
4255+
// bf16 scalar path promotes to f32, no dot2 variant
4256+
#define CREATE_MM_NODOT2(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID, REQSUBGROUPSIZE) \
42194257
if (device->mul_mat ## ID ## _l[TYPE]) \
42204258
ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _len, NAMELC ## F16ACC ## _data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, REQSUBGROUPSIZE > 0, REQSUBGROUPSIZE); \
42214259
if (device->mul_mat ## ID ## _m[TYPE]) \
@@ -4250,15 +4288,14 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
42504288
CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_f16, matmul_f16, wg_denoms, warptile, vk_mat_mat_push_constants, 3, , 0);
42514289
CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_f16_f32, matmul_f16_f32, wg_denoms, warptile, vk_mat_mat_push_constants, 3, , 0);
42524290

4253-
CREATE_MM(GGML_TYPE_BF16, pipeline_matmul_bf16, matmul_bf16, , wg_denoms, warptile, vk_mat_mat_push_constants, 3, , 0);
4291+
CREATE_MM_NODOT2(GGML_TYPE_BF16, pipeline_matmul_bf16, matmul_bf16, , wg_denoms, warptile, vk_mat_mat_push_constants, 3, , 0);
42544292

42554293
CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q1_0], matmul_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42564294
CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_0], matmul_q4_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42574295
CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_1], matmul_q4_1_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42584296
CREATE_MM2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q5_0], matmul_q5_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42594297
CREATE_MM2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q5_1], matmul_q5_1_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42604298
CREATE_MM2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q8_0], matmul_q8_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
4261-
42624299
CREATE_MM2(GGML_TYPE_Q2_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q2_K], matmul_q2_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42634300
CREATE_MM2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q3_K], matmul_q3_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
42644301
CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_K], matmul_q4_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0);
@@ -4298,8 +4335,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
42984335
CREATE_MM(GGML_TYPE_F32, pipeline_matmul_id_f32, matmul_id_subgroup_f32_f32, , wg_denoms, warptile_id, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size_16);
42994336
CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_id_f16, matmul_id_subgroup_f16, wg_denoms, warptile_id, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size_16);
43004337
CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_id_f16_f32, matmul_id_subgroup_f16_f32, wg_denoms, warptile_id, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size_16);
4301-
CREATE_MM(GGML_TYPE_BF16, pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile_id, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size_16);
4302-
4338+
CREATE_MM_NODOT2(GGML_TYPE_BF16, pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile_id, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size_16);
43034339
CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
43044340
CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
43054341
CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size);
@@ -4344,8 +4380,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
43444380
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);
43454381
CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_id_f16, matmul_id_f16, wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
43464382
CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_id_f16_f32, matmul_id_f16_f32, wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
4347-
CREATE_MM(GGML_TYPE_BF16, pipeline_matmul_id_bf16, matmul_id_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
4348-
4383+
CREATE_MM_NODOT2(GGML_TYPE_BF16, pipeline_matmul_id_bf16, matmul_id_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
43494384
CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_q1_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
43504385
CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_q4_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
43514386
CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_q4_1_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0);
@@ -4390,6 +4425,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
43904425
#undef CREATE_MM2
43914426
#undef CREATE_MMQ
43924427
#undef CREATE_MM
4428+
#undef CREATE_MM_NODOT2
43934429
} else {
43944430
// Create 6 variants, {s,m,l}x{unaligned,aligned}
43954431
#define CREATE_MM(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID, REQSUBGROUPSIZE) \
@@ -5453,6 +5489,7 @@ static vk_device ggml_vk_get_device(size_t idx) {
54535489
device->integer_dot_product = false;
54545490
device->shader_64b_indexing = false;
54555491
bool bfloat16_support = false;
5492+
bool dot2_f16_support = false;
54565493

54575494
for (const auto& properties : ext_props) {
54585495
if (strcmp("VK_KHR_maintenance4", properties.extensionName) == 0) {
@@ -5495,6 +5532,9 @@ static vk_device ggml_vk_get_device(size_t idx) {
54955532
!getenv("GGML_VK_DISABLE_BFLOAT16")) {
54965533
bfloat16_support = true;
54975534
#endif
5535+
} else if (strcmp("VK_VALVE_shader_mixed_float_dot_product", properties.extensionName) == 0 &&
5536+
!getenv("GGML_VK_DISABLE_DOT2")) {
5537+
dot2_f16_support = true;
54985538
} else if (strcmp("VK_KHR_pipeline_executable_properties", properties.extensionName) == 0) {
54995539
pipeline_executable_properties_support = true;
55005540
} else if (strcmp("VK_EXT_memory_priority", properties.extensionName) == 0 &&
@@ -5802,6 +5842,14 @@ static vk_device ggml_vk_get_device(size_t idx) {
58025842
device_extensions.push_back("VK_KHR_shader_integer_dot_product");
58035843
}
58045844

5845+
VkPhysicalDeviceShaderMixedFloatDotProductFeaturesVALVE dot2_features {};
5846+
dot2_features.sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_MIXED_FLOAT_DOT_PRODUCT_FEATURES_VALVE;
5847+
if (dot2_f16_support) {
5848+
last_struct->pNext = (VkBaseOutStructure *)&dot2_features;
5849+
last_struct = (VkBaseOutStructure *)&dot2_features;
5850+
device_extensions.push_back("VK_VALVE_shader_mixed_float_dot_product");
5851+
}
5852+
58055853
VkPhysicalDevicePipelineExecutablePropertiesFeaturesKHR pep_features {};
58065854
pep_features.sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_PIPELINE_EXECUTABLE_PROPERTIES_FEATURES_KHR;
58075855
if (pipeline_executable_properties_support) {
@@ -5836,6 +5884,8 @@ static vk_device ggml_vk_get_device(size_t idx) {
58365884
device->bf16 = false;
58375885
#endif
58385886

5887+
device->dot2_f16 = dot2_f16_support && dot2_features.shaderMixedFloatDotProductFloat16AccFloat32;
5888+
58395889
device->pipeline_robustness = pl_robustness_features.pipelineRobustness;
58405890

58415891
device->multi_add = vk12_props.shaderRoundingModeRTEFloat16 &&
@@ -6250,6 +6300,7 @@ static void ggml_vk_print_gpu_info(size_t idx) {
62506300
bool coopmat2_decode_vector_support = false;
62516301
bool integer_dot_product = false;
62526302
bool bfloat16_support = false;
6303+
bool dot2_f16_support = false;
62536304

62546305
for (auto properties : ext_props) {
62556306
if (strcmp("VK_KHR_16bit_storage", properties.extensionName) == 0) {
@@ -6279,6 +6330,9 @@ static void ggml_vk_print_gpu_info(size_t idx) {
62796330
!getenv("GGML_VK_DISABLE_BFLOAT16")) {
62806331
bfloat16_support = true;
62816332
#endif
6333+
} else if (strcmp("VK_VALVE_shader_mixed_float_dot_product", properties.extensionName) == 0 &&
6334+
!getenv("GGML_VK_DISABLE_DOT2")) {
6335+
dot2_f16_support = true;
62826336
}
62836337
}
62846338

@@ -6369,6 +6423,13 @@ static void ggml_vk_print_gpu_info(size_t idx) {
63696423
last_struct = (VkBaseOutStructure *)&coopmat2_decode_vector_features;
63706424
}
63716425

6426+
VkPhysicalDeviceShaderMixedFloatDotProductFeaturesVALVE dot2_features {};
6427+
dot2_features.sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_MIXED_FLOAT_DOT_PRODUCT_FEATURES_VALVE;
6428+
if (dot2_f16_support) {
6429+
last_struct->pNext = (VkBaseOutStructure *)&dot2_features;
6430+
last_struct = (VkBaseOutStructure *)&dot2_features;
6431+
}
6432+
63726433
vkGetPhysicalDeviceFeatures2(physical_device, &device_features2);
63736434

63746435
fp16 = fp16 && vk12_features.shaderFloat16;
@@ -6415,9 +6476,12 @@ static void ggml_vk_print_gpu_info(size_t idx) {
64156476
: coopmat_support ? "KHR_coopmat"
64166477
: "none";
64176478

6479+
bool dot2_f16 = dot2_f16_support && dot2_features.shaderMixedFloatDotProductFloat16AccFloat32;
6480+
const char *fp16_str = fp16 ? (dot2_f16 ? "dot2" : "1") : "0";
6481+
64186482
std::string device_name = props2.properties.deviceName.data();
6419-
GGML_LOG_DEBUG("ggml_vulkan: %zu = %s (%s) | uma: %d | fp16: %d | bf16: %d | warp size: %zu | shared memory: %d | int dot: %d | matrix cores: %s\n",
6420-
idx, device_name.c_str(), driver_props.driverName.data(), uma, fp16, bf16, subgroup_size,
6483+
GGML_LOG_DEBUG("ggml_vulkan: %zu = %s (%s) | uma: %d | fp16: %s | bf16: %d | warp size: %zu | shared memory: %d | int dot: %d | matrix cores: %s\n",
6484+
idx, device_name.c_str(), driver_props.driverName.data(), uma, fp16_str, bf16, subgroup_size,
64216485
props2.properties.limits.maxComputeSharedMemorySize, integer_dot_product, matrix_cores.c_str());
64226486

64236487
if (props2.properties.deviceType == vk::PhysicalDeviceType::eCpu) {
Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,27 @@
1+
#ifdef DOT2_F16
2+
#extension GL_EXT_spirv_intrinsics : require
3+
4+
spirv_instruction(extensions = ["SPV_VALVE_mixed_float_dot_product"],
5+
capabilities = [6912], id = 6916)
6+
float v_dot2_f32_f16(f16vec2 a, f16vec2 b, float acc);
7+
8+
ACC_TYPE dot_product(f16vec4 a, f16vec4 b, ACC_TYPE acc) {
9+
return ACC_TYPE(v_dot2_f32_f16(a.zw, b.zw, v_dot2_f32_f16(a.xy, b.xy, float(acc))));
10+
}
11+
12+
ACC_TYPE dot_product(f16vec2 a, f16vec2 b, ACC_TYPE acc) {
13+
return ACC_TYPE(v_dot2_f32_f16(a, b, float(acc)));
14+
}
15+
16+
#else
17+
18+
ACC_TYPE dot_product(FLOAT_TYPEV4 a, FLOAT_TYPEV4 b, ACC_TYPE acc) {
19+
return fma(ACC_TYPE(a.x), ACC_TYPE(b.x), fma(ACC_TYPE(a.y), ACC_TYPE(b.y),
20+
fma(ACC_TYPE(a.z), ACC_TYPE(b.z), fma(ACC_TYPE(a.w), ACC_TYPE(b.w), acc))));
21+
}
22+
23+
ACC_TYPE dot_product(FLOAT_TYPEV2 a, FLOAT_TYPEV2 b, ACC_TYPE acc) {
24+
return fma(ACC_TYPE(a.x), ACC_TYPE(b.x), fma(ACC_TYPE(a.y), ACC_TYPE(b.y), acc));
25+
}
26+
27+
#endif

ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
#extension GL_KHR_shader_subgroup_vote : enable
2222

2323
#include "types.glsl"
24+
#include "dot_product_funcs.glsl"
2425
#include "flash_attn_base.glsl"
2526
#include "flash_attn_dequant.glsl"
2627

@@ -318,7 +319,7 @@ void main() {
318319
K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c * cols_per_iter + col_tid) * k_stride / 4 + d * D_split + d_tid]);
319320
}
320321
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
321-
Sf[r][c] += dot(ACC_TYPEV4(Q_cache[r]), ACC_TYPEV4(K_Tf));
322+
Sf[r][c] = dot_product(Q_cache[r], K_Tf, Sf[r][c]);
322323
}
323324
}
324325
}
@@ -341,7 +342,7 @@ void main() {
341342
K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c * cols_per_iter + col_tid) * k_stride / 4 + d * D_split + d_tid]);
342343
}
343344
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
344-
Sf[r][c] += dot(ACC_TYPEV4(Qf[tile_row(r) * qf_stride + d * D_split + d_tid]), ACC_TYPEV4(K_Tf));
345+
Sf[r][c] = dot_product(Qf[tile_row(r) * qf_stride + d * D_split + d_tid], K_Tf, Sf[r][c]);
345346
}
346347
}
347348
}

0 commit comments

Comments
 (0)