diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 5db74d2307..9bf00029d1 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4276,11 +4276,18 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmq_int = { 128, 64, 64, 32, subgroup_size_8, 32, 2, itm_m, itn_m, itk_m, subgroup_size_8 }; s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, itm_s, itn_s, itk_s, subgroup_size_8 }; - // Coopmat int8 cm1 shader uses larger workgroups for better occupancy - // Force wave32 for VOPD on RDNA3+ - l_warptile_mmq_cm1_int = { 512, 128, 128, 32, 32, 32, 2, itm_l, itn_l, itk_l, 32 }; - m_warptile_mmq_cm1_int = { 128, 64, 64, 32, 32, 32, 2, itm_m, itn_m, itk_m, 32 }; - s_warptile_mmq_cm1_int = { 32, 32, 32, 32, 32, 32, 2, itm_s, itn_s, itk_s, 32 }; + // RDNA3.5 preferred wave32 here + const bool cm1_use_wave32 = device->vendor_id == VK_VENDOR_ID_AMD && + device->subgroup_size_control && + device->subgroup_min_size <= 32 && device->subgroup_max_size >= 32; + const uint32_t cm1_sg = cm1_use_wave32 ? 32 : device->subgroup_size; + const auto cm1_bs = [cm1_sg](uint32_t bm, uint32_t bn) { + return cm1_sg * (bm / std::min(cm1_sg, bm)) * (bn / 32); + }; + + l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm_l, itn_l, itk_l, cm1_sg }; + m_warptile_mmq_cm1_int = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg }; + s_warptile_mmq_cm1_int = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg }; // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 1, 4, 4, 1, subgroup_size_8 }; @@ -4744,11 +4751,11 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #define CREATE_MMQ(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true, 32); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, l_ ## WARPTILE, 1, false, true, cm1_sg); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true, 32); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, m_ ## WARPTILE, 1, false, true, cm1_sg); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true, 32); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, s_ ## WARPTILE, 1, false, true, cm1_sg); \ // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \