diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index ebbb412e55f..b42f092aa75 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -2464,8 +2464,18 @@ static void ggml_vk_load_shaders(vk_device& device) { // chip specific tuning if ((device->architecture == AMD_GCN) && (device->driver_id != vk::DriverId::eAmdProprietary)) { + l_warptile = { 256, 128, 128, 16, 16, 64, 2, tm_l, tn_l, tk_l, 16 }; + l_warptile_mmq = { 256, 128, 128, 32, 16, 64, 2, tm_l, tn_l, tk_l, 16 }; + l_warptile_mmq_int = { 256, 128, 128, 32, 16, 64, 2, 4, 4, 1, 16 }; + l_warptile_mmqid = { 256, 128, 128, 32, 16, 64, 2, tm_l, tn_l, tk_l, 16 }; + m_warptile_mmq = m_warptile_mmq_int = { 256, 64, 64, 32, 16, 16, 2, 2, 2, 1, 16 }; m_warptile_mmqid = { 256, 64, 64, 32, 16, 16, 2, 2, 2, 1, 16 }; + + s_warptile = { 128, 64, 64, 16, 16, 32, 2, tm_s, tn_s, tk_s, 16 }; + s_warptile_mmq = { 128, 64, 64, 32, 16, 32, 2, tm_s, tn_s, tk_s, 16 }; + s_warptile_mmq_int = { 128, 64, 64, 32, 16, 32, 2, 2, 1, 1, 16 }; + s_warptile_mmqid = { 128, 64, 64, 32, 16, 32, 2, tm_s, tn_s, tk_s, 16 }; } l_mmq_wg_denoms = l_wg_denoms = {128, 128, 1 };