diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 39b4cd35980..c96a12e4231 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -395,6 +395,7 @@ enum vk_device_architecture { AMD_RDNA1, AMD_RDNA2, AMD_RDNA3, + AMD_RDNA4, INTEL_XE1, INTEL_XE2, NVIDIA_PRE_TURING, @@ -410,6 +411,7 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& bool amd_shader_core_properties = false; bool integer_dot_product = false; bool subgroup_size_control = false; + bool shader_float8 = false; for (const auto& properties : ext_props) { if (strcmp("VK_AMD_shader_core_properties", properties.extensionName) == 0) { @@ -418,6 +420,8 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& integer_dot_product = true; } else if (strcmp("VK_EXT_subgroup_size_control", properties.extensionName) == 0) { subgroup_size_control = true; + } else if (strcmp("VK_EXT_shader_float8", properties.extensionName) == 0) { + shader_float8 = true; } } @@ -444,6 +448,9 @@ static vk_device_architecture get_device_architecture(const vk::PhysicalDevice& if (shader_core_props_amd.wavefrontsPerSimd == 20) { return vk_device_architecture::AMD_RDNA1; } + if (shader_float8) { + return vk_device_architecture::AMD_RDNA4; + } if (integer_dot_props.integerDotProduct4x8BitPackedMixedSignednessAccelerated) { return vk_device_architecture::AMD_RDNA3; } @@ -4070,6 +4077,8 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std: case GGML_TYPE_Q5_1: block_a_size = std430_size({{16, 4}, {4, 4}, {fp2_size, fp2_align}}); break; // qs[16/4] + qh + dm(vec2) case GGML_TYPE_Q8_0: block_a_size = std430_size({{32, 4}, {fp_size, fp_align}}); break; // qs[8] + dm case GGML_TYPE_MXFP4: block_a_size = std430_size({{32, 4}, {fp_size, fp_align}}); break; // qs[8] + d + case GGML_TYPE_IQ4_NL: block_a_size = std430_size({{32, 4}, {fp_size, fp_align}}); break; // qs[8] + d + case GGML_TYPE_NVFP4: block_a_size = std430_size({{32, 4}, {fp2_size, fp2_align}}); break; // qs[8] + d_scales(vec2) case GGML_TYPE_Q2_K: block_a_size = std430_size({{ 8, 4}, {2, 2}, {fp2_size, fp2_align}}); break; // qs[2] + scales(u8vec2) + dm(vec2) case GGML_TYPE_Q3_K: block_a_size = std430_size({{16, 4}, {fp2_size, fp2_align}}); break; // qs[4] + d_scales(vec2) case GGML_TYPE_Q4_K: block_a_size = std430_size({{16, 4}, {fp2_size, fp2_align}}); break; // qs[4] + dm(vec2) @@ -4103,6 +4112,66 @@ static bool ggml_vk_matmul_int_shmem_support(const vk_device& device, const std: return supported; } +static bool ggml_vk_matmul_cm1_int_shmem_support(const vk_device& device, const std::vector& warptile, bool mul_mat_id, ggml_type src0_type) { + + bool kscales2 = false; // two scale sets per block + bool has_dm = false; // d+m as vec2 + b-side sum + bool has_kvalues = false; + switch (src0_type) { + case GGML_TYPE_Q4_0: case GGML_TYPE_Q5_0: case GGML_TYPE_Q8_0: + break; + case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_1: + case GGML_TYPE_Q4_K: case GGML_TYPE_Q5_K: + has_dm = true; break; + case GGML_TYPE_IQ4_NL: case GGML_TYPE_MXFP4: + has_kvalues = true; break; + case GGML_TYPE_Q3_K: case GGML_TYPE_Q6_K: + kscales2 = true; break; + case GGML_TYPE_NVFP4: + kscales2 = true; has_kvalues = true; break; + default: + return false; + } + + const uint32_t BLOCK_SIZE = warptile[0]; + const uint32_t BM = warptile[1]; + const uint32_t BN = warptile[2]; + const uint32_t WARP = warptile[10]; + + const uint32_t BK = 32; + const uint32_t BK_STEP = mul_mat_id ? 2u : 4u; + const uint32_t QPITCH = BK_STEP * (BK / 4u) + 4u; + const uint32_t KSCALES = kscales2 ? 2u : 1u; + + uint32_t total = 0; + total += BM * QPITCH * (uint32_t)sizeof(uint32_t); // buf_a_qs + total += BN * QPITCH * (uint32_t)sizeof(uint32_t); // buf_b_qs + total += has_dm ? (BM * BK_STEP * 2u * (uint32_t)sizeof(float)) // buf_a_dm (vec2) + : (BM * BK_STEP * KSCALES * (uint32_t)sizeof(float)); // buf_a_d + total += BN * BK_STEP * (uint32_t)sizeof(float); // buf_b_d + if (has_dm) { + total += BN * BK_STEP * (uint32_t)sizeof(float); // buf_b_s + } + if (has_kvalues) { + total += 16u * (uint32_t)sizeof(int8_t); // cm1_kvalues[16] + } + if (src0_type == GGML_TYPE_NVFP4 && !device->ocp_fp4) { + total += 128u * (uint32_t)sizeof(float); // ue4m3_fp32_lut[128] + } + if (mul_mat_id) { + total += BN * 2u * (uint32_t)sizeof(uint16_t); // row_ids[BN] (u16vec2) + const uint32_t num_warps = BLOCK_SIZE / std::max(WARP, 1u); + total += num_warps * 4u * (uint32_t)sizeof(uint32_t); // ballots_sh[NUM_WARPS] (uvec4) + } + + const bool supported = total <= device->properties.limits.maxComputeSharedMemorySize; + + VK_LOG_DEBUG("ggml_vk_matmul_cm1_int_shmem_support(warptile=(" << warptile[0] << "," << warptile[1] << "," << warptile[2] << "), " + "mul_mat_id=" << mul_mat_id << ", src0_type=" << ggml_type_name(src0_type) << ", total=" << total << ", supported=" << supported); + + return supported; +} + struct GpuPipelineConfig { // GPU architecture identifier. // Example: vk_device_architecture::AMD_GCN @@ -4239,6 +4308,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_warptile_id, m_warptile_id, s_warptile_id, l_warptile_mmq, m_warptile_mmq, s_warptile_mmq, l_warptile_mmq_int, m_warptile_mmq_int, s_warptile_mmq_int, + l_warptile_mmq_cm1_int, m_warptile_mmq_cm1_int, s_warptile_mmq_cm1_int, + l_warptile_mmq_cm1_int_k, m_warptile_mmq_cm1_int_k, s_warptile_mmq_cm1_int_k, l_warptile_mmq_int_k, m_warptile_mmq_int_k, s_warptile_mmq_int_k, l_warptile_mmq_k, m_warptile_mmq_k, s_warptile_mmq_k, l_warptile_mmqid, m_warptile_mmqid, s_warptile_mmqid, @@ -4247,10 +4318,17 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { std::array l_wg_denoms, m_wg_denoms, s_wg_denoms, l_mmq_wg_denoms, m_mmq_wg_denoms, s_mmq_wg_denoms, l_mmq_wg_denoms_k, m_mmq_wg_denoms_k, s_mmq_wg_denoms_k, + l_mmq_cm1_wg_denoms_k, m_mmq_cm1_wg_denoms_k, s_mmq_cm1_wg_denoms_k, l_mmqid_wg_denoms, m_mmqid_wg_denoms, s_mmqid_wg_denoms; uint32_t l_align, m_align, s_align; + // 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; + vk_pipeline wait_pipeline; CompileTask claimed_task {}; bool has_claimed_task = false; @@ -4308,6 +4386,16 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t tk_m = device->coopmat_support ? device->coopmat_k : 1; const uint32_t tk_s = device->coopmat_support ? device->coopmat_k : 1; + const uint32_t itm_l = device->coopmat_int_support ? device->coopmat_int_m : 4; + const uint32_t itm_m = device->coopmat_int_support ? device->coopmat_int_m : 4; + const uint32_t itm_s = device->coopmat_int_support ? device->coopmat_int_m : 2; + const uint32_t itn_l = device->coopmat_int_support ? device->coopmat_int_n : 4; + const uint32_t itn_m = device->coopmat_int_support ? device->coopmat_int_n : 2; + const uint32_t itn_s = device->coopmat_int_support ? device->coopmat_int_n : 1; + const uint32_t itk_l = device->coopmat_int_support ? device->coopmat_int_k : 1; + const uint32_t itk_m = device->coopmat_int_support ? device->coopmat_int_k : 1; + const uint32_t itk_s = device->coopmat_int_support ? device->coopmat_int_k : 1; + const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32; l_warptile = { 128, 128, 128, 16, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 }; @@ -4319,9 +4407,25 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { s_warptile_mmq = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; // Integer MMQ has a smaller shared memory profile, but heavier register use - l_warptile_mmq_int = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, 4, 4, 1, mm_warp_8 }; - m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, 2, 2, 1, mm_warp_8 }; - s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, subgroup_size_8 }; + l_warptile_mmq_int = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, itm_l, itn_l, itk_l, mm_warp_8 }; + m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, itm_m, itn_m, itk_m, mm_warp_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 }; + + 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, (uint32_t)device->architecture }; + 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, (uint32_t)device->architecture }; + 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, (uint32_t)device->architecture }; + + l_warptile_mmq_cm1_int_k = { cm1_bs( 64, 128), 64, 128, 32, std::min(cm1_sg, 64u), 32, 2, itm_l, itn_l, itk_l, cm1_sg, (uint32_t)device->architecture }; + m_warptile_mmq_cm1_int_k = { cm1_bs( 64, 64), 64, 64, 32, std::min(cm1_sg, 64u), 32, 2, itm_m, itn_m, itk_m, cm1_sg, (uint32_t)device->architecture }; + s_warptile_mmq_cm1_int_k = { cm1_bs( 32, 32), 32, 32, 32, std::min(cm1_sg, 32u), 32, 2, itm_s, itn_s, itk_s, cm1_sg, (uint32_t)device->architecture }; + + l_mmq_cm1_wg_denoms_k = { l_warptile_mmq_cm1_int_k[1], l_warptile_mmq_cm1_int_k[2], 1 }; + m_mmq_cm1_wg_denoms_k = { m_warptile_mmq_cm1_int_k[1], m_warptile_mmq_cm1_int_k[2], 1 }; + s_mmq_cm1_wg_denoms_k = { s_warptile_mmq_cm1_int_k[1], s_warptile_mmq_cm1_int_k[2], 1 }; // K-quants use even more registers, mitigate by setting WMITER to 1 l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; @@ -4366,6 +4470,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_align = 64; s_align = 32; + const bool use_cm1_int = device->coopmat_int_support && + (device->architecture == AMD_RDNA3 || device->architecture == AMD_RDNA4); + for (uint32_t i = 0; i < GGML_TYPE_COUNT; ++i) { ggml_type t = (ggml_type)i; // Disable medium and large matrix multiplication if not enough shared memory is available @@ -4393,37 +4500,50 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { device->mul_mat_id_l[i] = false; } - // The q8_1 mmq path has its own (larger) shmem layout, check it separately. - // K-quants use the _int_k warptiles, others use _int. + // cm1 splits k-tiles on the KSCALES==2 types and shares tiles between dense/id. const bool is_k_quant = (t == GGML_TYPE_Q2_K || t == GGML_TYPE_Q3_K || t == GGML_TYPE_Q4_K || t == GGML_TYPE_Q5_K || t == GGML_TYPE_Q6_K); - const auto & s_int = is_k_quant ? s_warptile_mmq_int_k : s_warptile_mmq_int; - const auto & m_int = is_k_quant ? m_warptile_mmq_int_k : m_warptile_mmq_int; - const auto & l_int = is_k_quant ? l_warptile_mmq_int_k : l_warptile_mmq_int; - const auto & s_intid = is_k_quant ? s_warptile_mmqid_int_k : s_warptile_mmqid_int; - const auto & m_intid = is_k_quant ? m_warptile_mmqid_int_k : m_warptile_mmqid_int; - const auto & l_intid = is_k_quant ? l_warptile_mmqid_int_k : l_warptile_mmqid_int; - - if (!ggml_vk_matmul_int_shmem_support(device, s_int, false, t)) { + const bool cm1_k_tile = (t == GGML_TYPE_Q3_K || t == GGML_TYPE_Q6_K || + t == GGML_TYPE_NVFP4); + + const auto & s_int = use_cm1_int ? (cm1_k_tile ? s_warptile_mmq_cm1_int_k : s_warptile_mmq_cm1_int) + : (is_k_quant ? s_warptile_mmq_int_k : s_warptile_mmq_int); + const auto & m_int = use_cm1_int ? (cm1_k_tile ? m_warptile_mmq_cm1_int_k : m_warptile_mmq_cm1_int) + : (is_k_quant ? m_warptile_mmq_int_k : m_warptile_mmq_int); + const auto & l_int = use_cm1_int ? (cm1_k_tile ? l_warptile_mmq_cm1_int_k : l_warptile_mmq_cm1_int) + : (is_k_quant ? l_warptile_mmq_int_k : l_warptile_mmq_int); + const auto & s_intid = use_cm1_int ? (cm1_k_tile ? s_warptile_mmq_cm1_int_k : s_warptile_mmq_cm1_int) + : (is_k_quant ? s_warptile_mmqid_int_k : s_warptile_mmqid_int); + const auto & m_intid = use_cm1_int ? (cm1_k_tile ? m_warptile_mmq_cm1_int_k : m_warptile_mmq_cm1_int) + : (is_k_quant ? m_warptile_mmqid_int_k : m_warptile_mmqid_int); + const auto & l_intid = use_cm1_int ? (cm1_k_tile ? l_warptile_mmq_cm1_int_k : l_warptile_mmq_cm1_int) + : (is_k_quant ? l_warptile_mmqid_int_k : l_warptile_mmqid_int); + + const auto int_shmem_support = [&](const std::vector& wt, bool id) { + return use_cm1_int ? ggml_vk_matmul_cm1_int_shmem_support(device, wt, id, t) + : ggml_vk_matmul_int_shmem_support(device, wt, id, t); + }; + + if (!int_shmem_support(s_int, false)) { device->mul_mat_s_int[i] = false; device->mul_mat_m_int[i] = false; device->mul_mat_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, m_int, false, t)) { + } else if (!int_shmem_support(m_int, false)) { device->mul_mat_m_int[i] = false; device->mul_mat_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, l_int, false, t)) { + } else if (!int_shmem_support(l_int, false)) { device->mul_mat_l_int[i] = false; } - if (!ggml_vk_matmul_int_shmem_support(device, s_intid, true, t)) { + if (!int_shmem_support(s_intid, true)) { device->mul_mat_id_s_int[i] = false; device->mul_mat_id_m_int[i] = false; device->mul_mat_id_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, m_intid, true, t)) { + } else if (!int_shmem_support(m_intid, true)) { device->mul_mat_id_m_int[i] = false; device->mul_mat_id_l_int[i] = false; - } else if (!ggml_vk_matmul_int_shmem_support(device, l_intid, true, t)) { + } else if (!int_shmem_support(l_intid, true)) { device->mul_mat_id_l_int[i] = false; } } @@ -4783,6 +4903,14 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (device->mul_mat ## ID ## _s[TYPE]) \ ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true); \ +#define CREATE_MMQ(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + if (device->mul_mat ## ID ## _l_int[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, cm1_sg); \ + if (device->mul_mat ## ID ## _m_int[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, cm1_sg); \ + if (device->mul_mat ## ID ## _s_int[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, cm1_sg); \ + // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->coopmat_acc_f16_support) { \ @@ -4792,6 +4920,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM(TYPE, PIPELINE_NAME . f32acc, NAMELC, , WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ } \ +#define CREATE_MMQ2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + CREATE_MMQ(TYPE, PIPELINE_NAME . f16acc, NAMELC, _f16acc, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + CREATE_MMQ(TYPE, PIPELINE_NAME . f32acc, NAMELC, , WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + CREATE_MM(GGML_TYPE_F32, pipeline_matmul_f32, matmul_f32_f32, , wg_denoms, warptile, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_F32, pipeline_matmul_f32_f16, matmul_f32_f16, , wg_denoms, warptile, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_F16, pipeline_matmul_f16, matmul_f16, wg_denoms, warptile, vk_mat_mat_push_constants, 3, ); @@ -4837,6 +4969,37 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { 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, ); } + // Some quants are not performant on RDNA4, those fall back to FP16 matmul + const bool rdna3 = device->architecture == AMD_RDNA3; + const bool rdna4 = device->architecture == AMD_RDNA4; + if (device->coopmat_int_support && (rdna3 || rdna4)) { + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_0], matmul_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_1], matmul_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_0], matmul_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_1], matmul_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q8_0], matmul_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_IQ4_NL], matmul_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_MXFP4], matmul_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); + CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q3_K], matmul_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q4_K], matmul_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q5_K], matmul_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_push_constants, 3, ); } + CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_Q6_K], matmul_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_q8_1[GGML_TYPE_NVFP4], matmul_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_push_constants, 3, ); } + + CREATE_MMQ2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_q8_1, mmq_wg_denoms, warptile_mmq_cm1_int, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MMQ2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + if (!rdna4) { CREATE_MMQ2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_q8_1, mmq_cm1_wg_denoms_k, warptile_mmq_cm1_int_k, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } + } + GGML_ASSERT(device->subgroup_ballot); CREATE_MM(GGML_TYPE_F32, pipeline_matmul_id_f32, matmul_id_subgroup_f32_f32, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4880,6 +5043,8 @@ 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); } +#undef CREATE_MMQ2 +#undef CREATE_MMQ #undef CREATE_MM2 #undef CREATE_MM } else @@ -7831,10 +7996,19 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte assert(src1_type == GGML_TYPE_F16); return prec == GGML_PREC_DEFAULT ? ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f32acc; } + + vk_matmul_pipeline pipelines; if (ctx->device->coopmat_support) { - return (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + pipelines = (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + } else { + pipelines = (ctx->device->fp16 && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; } - return (ctx->device->fp16 && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + + if (pipelines->is_empty()) { + return nullptr; + } + + return pipelines; } static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec(ggml_backend_vk_context * ctx, ggml_type a_type, ggml_type b_type, uint32_t num_cols, uint32_t m, uint32_t k) { @@ -9294,7 +9468,8 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub const bool y_f32_kernel = src1->type == GGML_TYPE_F32 && !y_non_contig; - bool quantize_y = ctx->device->integer_dot_product && src1->type == GGML_TYPE_F32 && ggml_is_contiguous(src1) && !y_non_contig && (ne11 * ne10) % 4 == 0; + bool quantize_y = (ctx->device->integer_dot_product || ctx->device->coopmat_int_support) && + src1->type == GGML_TYPE_F32 && ggml_is_contiguous(src1) && !y_non_contig && (ne11 * ne10) % 4 == 0; // Check for mmq first vk_matmul_pipeline mmp = quantize_y ? ggml_vk_get_mul_mat_mat_pipeline(ctx, src0->type, GGML_TYPE_Q8_1, (ggml_prec)dst->op_params[0]) : nullptr; @@ -19273,7 +19448,7 @@ static bool ggml_vk_khr_cooperative_matrix_support(const vk::PhysicalDevicePrope case VK_VENDOR_ID_AMD: if (driver_props.driverID == vk::DriverId::eAmdProprietary || driver_props.driverID == vk::DriverId::eAmdOpenSource) { // Workaround for AMD proprietary driver reporting support on all GPUs - return arch == vk_device_architecture::AMD_RDNA3; + return arch == vk_device_architecture::AMD_RDNA3 || arch == vk_device_architecture::AMD_RDNA4; } return true; default: diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp new file mode 100644 index 00000000000..eccd69869b7 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp @@ -0,0 +1,506 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : enable +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +#extension GL_KHR_shader_subgroup_basic : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : enable + +#if defined(MUL_MAT_ID_USE_SUBGROUPS) +#extension GL_KHR_shader_subgroup_ballot : enable +#endif + +#ifdef MUL_MAT_ID +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#endif + +#include "types.glsl" + +#if defined(DATA_A_Q3_K) || defined(DATA_A_Q6_K) || defined(DATA_A_NVFP4) +#define KSCALES 2 +#else +#define KSCALES 1 +#endif + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; +#if defined(A_TYPE_PACKED16) +layout (binding = 0) readonly buffer A_PACKED16 {A_TYPE_PACKED16 data_a_packed16[];}; +#endif +#if defined(A_TYPE_PACKED32) +layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32[];}; +#endif +layout (binding = 1) readonly buffer B {block_q8_1_x4_packed128 data_b[];}; +layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; + +#ifdef MUL_MAT_ID +layout (binding = 3) readonly buffer IDS {int data_ids[];}; +layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +#endif + +layout (push_constant) uniform parameter +{ + uint M; + uint N; + uint K; + uint stride_a; + uint stride_b; + uint stride_d; + + uint batch_stride_a; + uint batch_stride_b; + uint batch_stride_d; + +#ifdef MUL_MAT_ID + uint nei0; + uint nei1; + uint nbi1; + uint ne11; + uint n_experts; + uint hoist_row_ids; +#else + uint base_work_group_z; + uint num_batches; + uint k_split; + uint ne02; + uint ne12; + uint broadcast2; + uint broadcast3; +#endif +} p; + +layout (constant_id = 0) const uint BLOCK_SIZE = 256; +layout (constant_id = 1) const uint BM = 128; +layout (constant_id = 2) const uint BN = 128; +// layout (constant_id = 3) const uint BK = 32; +layout (constant_id = 4) const uint WM = 64; +layout (constant_id = 5) const uint WN = 32; +layout (constant_id = 7) const uint TM = 16; +layout (constant_id = 8) const uint TN = 16; +layout (constant_id = 9) const uint TK = 16; +layout (constant_id = 10) const uint WARP = 32; +layout (constant_id = 11) const uint DEVICE_ARCH = 0; // vk_device_architecture (ggml-vulkan.cpp) +#define VK_ARCH_AMD_RDNA4 5u + +#define BK 32 +#ifdef MUL_MAT_ID +#define BK_STEP 2 +#else +#define BK_STEP 4 +#endif +#define GROUP_A_BUDGET (16u * 1024u * 1024u) + +const uint QPITCH = BK_STEP * (BK / 4) + 4; + +shared uint32_t buf_a_qs[BM * QPITCH]; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) +shared vec2 buf_a_dm[BM * BK_STEP]; // .x = d, .y = m +#else +shared float buf_a_d[BM * BK_STEP * KSCALES]; +#endif + +shared uint32_t buf_b_qs[BN * QPITCH]; +shared float buf_b_d[BN * BK_STEP]; + +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) +shared float buf_b_s[BN * BK_STEP]; +#endif + +#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4) || defined(DATA_A_NVFP4) +shared int8_t cm1_kvalues[16]; +#endif + +#if defined(DATA_A_QUANT_K) || defined(DATA_A_NVFP4) +#define LOAD_VEC_A 8 +#else +#define LOAD_VEC_A (4 * QUANT_R) +#endif +#define LOAD_VEC_B 16 + +const uint CM_ELEMS = (TM * TN) / WARP; +#define ACC_BIAS_BITS 0x4B400000 +#define ACC_BIAS_F 12582912.0f +const bool USE_MAGIC_BIAS = WARP != 32; + +// Accumulator row for element e: RDNA4 blocked, RDNA3/3.5 interleaved. +uint cm_elem_row(uint e) { + const uint row_half = gl_SubgroupInvocationID / TN; + return (DEVICE_ARCH == VK_ARCH_AMD_RDNA4) ? (e + row_half * CM_ELEMS) : (row_half + 2u * e); +} + +// min_term = asymmetric-quant min*b_sum correction (0 for symmetric types). +ACC_TYPE cm1_accumulate(ACC_TYPE prev, int acc_e, float scale_a, float nbias_a, float scale_b, float min_term) { + if (USE_MAGIC_BIAS) { + const float t = fma(intBitsToFloat(acc_e), scale_a, nbias_a); + return ACC_TYPE(fma(t, scale_b, float(prev) + min_term)); + } + return prev + ACC_TYPE(fma(float(acc_e) * scale_a, scale_b, min_term)); +} + +#ifdef MUL_MAT_ID +#define NUM_WARPS (BLOCK_SIZE / WARP) +#include "mul_mm_id_funcs.glsl" +#endif + +#include "mul_mmq_cm1_funcs.glsl" + +void main() { +#if defined(DATA_A_IQ4_NL) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex]; + } + barrier(); +#elif defined(DATA_A_MXFP4) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex]; + } + barrier(); +#elif defined(DATA_A_NVFP4) + if (gl_LocalInvocationIndex < 16u) { + cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex]; + } +#if !defined(USE_OCP_FP4) + for (uint i = gl_LocalInvocationIndex; i < 128u; i += BLOCK_SIZE) { + ue4m3_fp32_lut[i] = ue4m3_to_fp32_build(i); + } +#endif + barrier(); +#endif + + const uint blocks_m = (p.M + BM - 1) / BM; + const uint ik = gl_WorkGroupID.x / blocks_m; + +#ifdef MUL_MAT_ID + const uint ic = gl_WorkGroupID.y; + const uint ir = gl_WorkGroupID.x % blocks_m; + const uint expert_idx = gl_WorkGroupID.z; + if (ic * BN >= data_expert_count[expert_idx]) { + return; + } +#else + // L2-friendly workgroup scheduling + const uint blocks_n = (p.N + BN - 1) / BN; + const uint a_panel_bytes = BM * p.K + (BM * p.K) / 16; + const uint group_m = clamp(GROUP_A_BUDGET / max(a_panel_bytes, 1u), 1u, min(blocks_m, 32u)); + const uint tiles_per_group = group_m * blocks_n; + const uint lin = gl_WorkGroupID.y * blocks_m + (gl_WorkGroupID.x % blocks_m); + const uint group_id = lin / tiles_per_group; + const uint first_m = group_id * group_m; + const uint gsize = min(blocks_m - first_m, group_m); + const uint in_group = lin - group_id * tiles_per_group; + const uint ir = first_m + in_group % gsize; + const uint ic = in_group / gsize; +#endif + +#ifndef MUL_MAT_ID + const uint batch_idx = gl_WorkGroupID.z + p.base_work_group_z; + + const uint i13 = batch_idx / p.ne12; + const uint i12 = batch_idx % p.ne12; + + const uint i03 = i13 / p.broadcast3; + const uint i02 = i12 / p.broadcast2; + + const uint batch_idx_a = i03 * p.ne02 + i02; +#endif + + const uint warp_i = gl_SubgroupID; + + const uint cms_per_row = WM / TM; + const uint cms_per_col = WN / TN; + + const uint warp_r = warp_i % (BM / WM); + const uint warp_c = warp_i / (BM / WM); + + const uint elem_col0 = gl_SubgroupInvocationID % TN; + + const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A); + const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A); + const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B); + const uint loadc_b = gl_LocalInvocationID.x / (BK / LOAD_VEC_B); + + const uint loadstride_a = BLOCK_SIZE * LOAD_VEC_A / BK; + const uint loadstride_b = BLOCK_SIZE * LOAD_VEC_B / BK; + +#ifdef MUL_MAT_ID + if (p.hoist_row_ids != 0) { + load_row_ids_hoisted(expert_idx, ic); + } else { +#ifdef MUL_MAT_ID_USE_SUBGROUPS + if (bitCount(p.nei0) == 1) { + load_row_ids(expert_idx, true, ic); + } else { + load_row_ids(expert_idx, false, ic); + } +#else + _ne1 = 0; + for (uint ii1 = 0; ii1 < p.nei1 && _ne1 < (ic + 1) * BN; ii1++) { + for (uint ii0 = 0; ii0 < p.nei0 && _ne1 < (ic + 1) * BN; ii0++) { + if (data_ids[ii1*p.nbi1 + ii0] == expert_idx) { + if (_ne1 >= ic * BN) { + row_ids[_ne1 - ic * BN] = u16vec2(ii0, ii1); + } + _ne1++; + } + } + } + + barrier(); +#endif + } + + if (ic * BN >= _ne1) return; +#endif + +#ifdef MUL_MAT_ID + const uint start_k = 0; + const uint end_k = p.K; +#else + const uint start_k = ik * p.k_split; + const uint end_k = min(p.K, (ik + 1) * p.k_split); +#endif + + uint pos_a_ib = +#ifdef MUL_MAT_ID + expert_idx * (p.batch_stride_a / BK) + +#else + batch_idx_a * (p.batch_stride_a / BK) + +#endif + (ir * BM * p.stride_a + start_k) / BK; +#ifdef MUL_MAT_ID + uint pos_b_ib = 0; +#else + uint pos_b_ib = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / BK; +#endif + + ACC_TYPE sums[cms_per_row * cms_per_col * CM_ELEMS]; + [[unroll]] for (uint i = 0; i < cms_per_row * cms_per_col * CM_ELEMS; i++) { + sums[i] = ACC_TYPE(0.0); + } + + // Double-buffering: prefetch registers + const uint A_LOADS = (BM + loadstride_a - 1) / loadstride_a; + const uint B_LOADS = (BN + loadstride_b - 1) / loadstride_b; + + block_a_prefetch pre_a[A_LOADS * BK_STEP]; + block_b_prefetch pre_b[B_LOADS * BK_STEP]; + + if (start_k < end_k) { + PREFETCH_BLOCK(start_k) + } + + const uint a_row0 = warp_r * WM; + const uint b_col0 = warp_c * WN; +#ifdef MUL_MAT_ID + const bool active_col_tile = ic * BN + b_col0 < _ne1; +#else + const bool active_col_tile = ic * BN + b_col0 < p.N; +#endif + + barrier(); + + for (uint block = start_k; block < end_k; block += BK * BK_STEP) { + STORE_BLOCK_TO_LDS(block) + + barrier(); + + pos_a_ib += BK_STEP; + pos_b_ib += BK_STEP; + + const uint next_block = block + BK * BK_STEP; + if (next_block < end_k) { + PREFETCH_BLOCK(next_block) + } + + if (active_col_tile) { + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { + const uint K_SUB = BK / TK; + +#if KSCALES == 2 + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + coopmat cache_a; + coopMatLoad(cache_a, buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + + float scale_a[CM_ELEMS]; + float nbias_a[CM_ELEMS]; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + cm_elem_row(e)]; + if (USE_MAGIC_BIAS) { + nbias_a[e] = -ACC_BIAS_F * scale_a[e]; + } + } + + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat cache_b; + coopMatLoad(cache_b, buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + + const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0]; + + coopmat acc = + coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); + acc = coopMatMulAdd(cache_a, cache_b, acc); + + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + sums[tile_idx * CM_ELEMS + e] = cm1_accumulate( + sums[tile_idx * CM_ELEMS + e], int(acc[e]), + scale_a[e], nbias_a[e], scale_b_v, 0.0); + } + } + } + } +#elif defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + // Preload all A/B fragments up front (ILP). + coopmat cache_a[cms_per_row * K_SUB]; + coopmat cache_b[cms_per_col * K_SUB]; + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + } + } + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + } + } + + float scale_b[cms_per_col]; + float bs[cms_per_col]; + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0]; + bs[c] = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col0]); + } + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + float scale_a[CM_ELEMS]; + float nbias_a[CM_ELEMS]; + float ma[CM_ELEMS]; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + cm_elem_row(e)]; + scale_a[e] = dm.x; + if (USE_MAGIC_BIAS) { + nbias_a[e] = -ACC_BIAS_F * scale_a[e]; + } + ma[e] = dm.y; + } + + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat acc = + coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); + + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); + } + + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + sums[tile_idx * CM_ELEMS + e] = cm1_accumulate( + sums[tile_idx * CM_ELEMS + e], int(acc[e]), + scale_a[e], nbias_a[e], scale_b[c], ma[e] * bs[c]); + } + } + } +#else + // Preload all A/B fragments up front (ILP). + coopmat cache_a[cms_per_row * K_SUB]; + coopmat cache_b[cms_per_col * K_SUB]; + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor); + } + } + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor); + } + } + + float scale_b[cms_per_col]; + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col0]; + } + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + float scale_a[CM_ELEMS]; + float nbias_a[CM_ELEMS]; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + cm_elem_row(e)]; + if (USE_MAGIC_BIAS) { + nbias_a[e] = -ACC_BIAS_F * scale_a[e]; + } + } + + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + coopmat acc = + coopmat( + USE_MAGIC_BIAS ? ACC_BIAS_BITS : 0); + + [[unroll]] for (uint h = 0; h < K_SUB; h++) { + acc = coopMatMulAdd(cache_a[r * K_SUB + h], cache_b[c * K_SUB + h], acc); + } + + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + sums[tile_idx * CM_ELEMS + e] = cm1_accumulate( + sums[tile_idx * CM_ELEMS + e], int(acc[e]), + scale_a[e], nbias_a[e], scale_b[c], 0.0); + } + } + } +#endif // KSCALES + } + } + + barrier(); + } + +#undef PREFETCH_BLOCK +#undef STORE_BLOCK_TO_LDS +#undef B_IB_CALC + + const uint dr = ir * BM + a_row0; + const uint dc = ic * BN + b_col0; + +#ifdef MUL_MAT_ID + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + const uint col_i = dc + c * TN + elem_col0; + if (col_i >= _ne1) continue; + + const uint row_g = dr + r * TM + cm_elem_row(e); + if (row_g >= p.M) continue; + + const u16vec2 row_idx = row_ids[col_i - ic * BN]; + const uint store_offset = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + row_g; + data_d[store_offset] = D_TYPE(sums[tile_idx * CM_ELEMS + e]); + } + } + } +#else + const uint offsets = batch_idx * p.batch_stride_d + ik * p.batch_stride_d * p.num_batches; + + [[unroll]] for (uint r = 0; r < cms_per_row; r++) { + [[unroll]] for (uint c = 0; c < cms_per_col; c++) { + const uint tile_idx = r * cms_per_col + c; + [[unroll]] for (uint e = 0; e < CM_ELEMS; e++) { + const uint row_g = dr + r * TM + cm_elem_row(e); + const uint col_g = dc + c * TN + elem_col0; + if (row_g < p.M && col_g < p.N) { + data_d[offsets + col_g * p.stride_d + row_g] = D_TYPE(sums[tile_idx * CM_ELEMS + e]); + } + } + } + } +#endif // MUL_MAT_ID +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl new file mode 100644 index 00000000000..3b07934f989 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl @@ -0,0 +1,558 @@ +// Per-quant-type data structures and functions for the cm1 int8 coopmat path. +// Each quant type defines: +// struct block_a_prefetch — register data for one A-block per thread +// block_a_load() — load from global memory into a block_a_prefetch +// block_a_to_shmem() — unpack and write to shared memory + +#if defined(DATA_A_Q4_0) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; + hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_Q4_1) + +struct block_a_prefetch { + uint32_t qs; + f16vec2 dm; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = data_a_packed32[ib].qs[loadr]; + blk.dm = data_a_packed32[ib].dm; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned nibbles; the -8 offset is absorbed by the min term. + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); + } +} + +#elif defined(DATA_A_Q5_0) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; + uint32_t qh; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + blk.qh = pack32(u16vec2(data_a_packed16[ib].qh[0], data_a_packed16[ib].qh[1])); + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 |= ((blk.qh >> (4u * loadr )) & 0xFu) * 0x02040810u & 0x10101010u; + hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; + lo4 = ((lo4 | 0x80808080) - 0x10101010) ^ 0x80808080; + hi4 = ((hi4 | 0x80808080) - 0x10101010) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_Q5_1) + +struct block_a_prefetch { + uint32_t qs; + f16vec2 dm; + uint32_t qh; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = data_a_packed32[ib].qs[loadr]; + blk.dm = data_a_packed32[ib].dm; + blk.qh = data_a_packed32[ib].qh; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned 5-bit values; the -16 offset is absorbed by the min term. + uint32_t lo4 = blk.qs & 0x0F0F0F0F; + uint32_t hi4 = (blk.qs >> 4) & 0x0F0F0F0F; + lo4 |= ((blk.qh >> (4u * loadr )) & 0xFu) * 0x02040810u & 0x10101010u; + hi4 |= ((blk.qh >> (4u * loadr + 16u )) & 0xFu) * 0x02040810u & 0x10101010u; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = lo4; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4; + + if (loadr == 0) { + buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y)); + } +} + +#elif defined(DATA_A_Q8_0) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr] = blk.qs; + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_IQ4_NL) + +struct block_a_prefetch { + uint32_t qs; + float16_t d; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u16vec2(data_a_packed16[ib].qs[loadr * 2], + data_a_packed16[ib].qs[loadr * 2 + 1])); + blk.d = data_a_packed16[ib].d; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); + const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = float(blk.d); + } +} + +#elif defined(DATA_A_MXFP4) + +struct block_a_prefetch { + uint32_t qs; + uint8_t e; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + blk.qs = pack32(u8vec4(data_a[ib].qs[loadr * 4], + data_a[ib].qs[loadr * 4 + 1], + data_a[ib].qs[loadr * 4 + 2], + data_a[ib].qs[loadr * 4 + 3])); + blk.e = data_a[ib].e; + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); + const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] = + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); + + if (loadr == 0) { + buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5; + } +} + +// LOAD_VEC_A=8 for k-quants and NVFP4: loadr has 4 positions, each writes 2 uint32 + +#elif defined(DATA_A_Q4_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint qs_base = (sub >> 1) * 8; + + uint32_t raw0 = data_a_packed32[ib_k].qs[qs_base + loadr * 2]; + uint32_t raw1 = data_a_packed32[ib_k].qs[qs_base + loadr * 2 + 1]; + if ((sub & 1u) != 0u) { + blk.qs0 = (raw0 >> 4) & 0x0F0F0F0F; + blk.qs1 = (raw1 >> 4) & 0x0F0F0F0F; + } else { + blk.qs0 = raw0 & 0x0F0F0F0F; + blk.qs1 = raw1 & 0x0F0F0F0F; + } + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned nibbles (blk.qs already masked); no -8 recentering needed. + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = blk.qs0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = blk.qs1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + const uint j = sub & 3u; + const uint s_j = uint(data_a[ib_k].scales[j]); + const uint s_j4 = uint(data_a[ib_k].scales[j + 4]); + const uint s_j8 = uint(data_a[ib_k].scales[j + 8]); + const uint sc_val = (sub < 4) ? (s_j & 0x3Fu) : ((s_j8 & 0x0Fu) | ((s_j >> 6) << 4)); + const uint mn_val = (sub < 4) ? (s_j4 & 0x3Fu) : ((s_j8 >> 4) | ((s_j4 >> 6) << 4)); + vec2 dm = vec2(data_a_packed32[ib_k].dm); + float d_scaled = dm.x * float(sc_val); + buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, -(dm.y * float(mn_val))); + } +} + +#elif defined(DATA_A_Q5_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint32_t qh0; + uint32_t qh1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint qs_base = (sub >> 1) * 8; + + uint32_t raw0 = data_a_packed32[ib_k].qs[qs_base + loadr * 2]; + uint32_t raw1 = data_a_packed32[ib_k].qs[qs_base + loadr * 2 + 1]; + if ((sub & 1u) != 0u) { + blk.qs0 = (raw0 >> 4) & 0x0F0F0F0F; + blk.qs1 = (raw1 >> 4) & 0x0F0F0F0F; + } else { + blk.qs0 = raw0 & 0x0F0F0F0F; + blk.qs1 = raw1 & 0x0F0F0F0F; + } + blk.qh0 = ((data_a_packed32[ib_k].qh[loadr * 2 ] >> sub) & 0x01010101) << 4; + blk.qh1 = ((data_a_packed32[ib_k].qh[loadr * 2 + 1] >> sub) & 0x01010101) << 4; + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + // Store raw unsigned 5-bit values (qs nibble | qh bit); no -16 recentering needed. + uint32_t v0 = blk.qs0 | blk.qh0; + uint32_t v1 = blk.qs1 | blk.qh1; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + const uint j = sub & 3u; + const uint s_j = uint(data_a[ib_k].scales[j]); + const uint s_j4 = uint(data_a[ib_k].scales[j + 4]); + const uint s_j8 = uint(data_a[ib_k].scales[j + 8]); + const uint sc_val = (sub < 4) ? (s_j & 0x3Fu) : ((s_j8 & 0x0Fu) | ((s_j >> 6) << 4)); + const uint mn_val = (sub < 4) ? (s_j4 & 0x3Fu) : ((s_j8 >> 4) | ((s_j4 >> 6) << 4)); + vec2 dm = vec2(data_a_packed32[ib_k].dm); + float d_scaled = dm.x * float(sc_val); + buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, -(dm.y * float(mn_val))); + } +} + +#elif defined(DATA_A_Q6_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint g = sub / 4; + const uint j = sub % 4; + + const uint ql_u16 = g * 32 + (j & 1) * 16 + loadr * 4; + const uint qh_u16 = g * 16 + loadr * 4; + const uint qh_shift = j * 2; + + uint32_t ql0 = pack32(u16vec2(data_a_packed16[ib_k].ql[ql_u16 ], + data_a_packed16[ib_k].ql[ql_u16 + 1])); + uint32_t ql1 = pack32(u16vec2(data_a_packed16[ib_k].ql[ql_u16 + 2], + data_a_packed16[ib_k].ql[ql_u16 + 3])); + if (j >= 2) { + ql0 = (ql0 >> 4) & 0x0F0F0F0F; + ql1 = (ql1 >> 4) & 0x0F0F0F0F; + } else { + ql0 = ql0 & 0x0F0F0F0F; + ql1 = ql1 & 0x0F0F0F0F; + } + + uint32_t qh0 = pack32(u16vec2(data_a_packed16[ib_k].qh[qh_u16 ], + data_a_packed16[ib_k].qh[qh_u16 + 1])); + uint32_t qh1 = pack32(u16vec2(data_a_packed16[ib_k].qh[qh_u16 + 2], + data_a_packed16[ib_k].qh[qh_u16 + 3])); + + blk.qs0 = ql0 | (((qh0 >> qh_shift) & 0x03030303) << 4); + blk.qs1 = ql1 | (((qh1 >> qh_shift) & 0x03030303) << 4); + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x20202020) ^ 0x80808080; + uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x20202020) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + i8vec2 sc = unpack8(int32_t(int16_t(data_a_packed16[ib_k].scales[sub]))).xy; + buf_a_d[(ks * KSCALES ) * BM + buf_ib] = float(data_a_packed16[ib_k].d) * float(sc.x); + buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = float(data_a_packed16[ib_k].d) * float(sc.y); + } +} + +#elif defined(DATA_A_Q3_K) + +struct block_a_prefetch { + uint32_t qs0; + uint32_t qs1; + uint ib; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 8; + const uint sub = ib % 8; + const uint g = sub / 4; + const uint j = sub % 4; + const uint qs_shift = j * 2; + const uint hm_bit = j + g * 4; + + const uint qs_u16 = g * 16 + loadr * 4; + uint32_t qs0 = pack32(u16vec2(data_a_packed16[ib_k].qs[qs_u16 ], + data_a_packed16[ib_k].qs[qs_u16 + 1])); + uint32_t qs1 = pack32(u16vec2(data_a_packed16[ib_k].qs[qs_u16 + 2], + data_a_packed16[ib_k].qs[qs_u16 + 3])); + + const uint hm_u16 = loadr * 4; + uint32_t hm0 = pack32(u16vec2(data_a_packed16[ib_k].hmask[hm_u16 ], + data_a_packed16[ib_k].hmask[hm_u16 + 1])); + uint32_t hm1 = pack32(u16vec2(data_a_packed16[ib_k].hmask[hm_u16 + 2], + data_a_packed16[ib_k].hmask[hm_u16 + 3])); + + blk.qs0 = ((qs0 >> qs_shift) & 0x03030303) | (((hm0 >> hm_bit) & 0x01010101) << 2); + blk.qs1 = ((qs1 >> qs_shift) & 0x03030303) | (((hm1 >> hm_bit) & 0x01010101) << 2); + blk.ib = ib; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + uint32_t v0 = ((blk.qs0 | 0x80808080) - 0x04040404) ^ 0x80808080; + uint32_t v1 = ((blk.qs1 | 0x80808080) - 0x04040404) ^ 0x80808080; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 ] = v0; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr * 2 + 1] = v1; + + if (loadr == 0) { + const uint ib_k = blk.ib / 8; + const uint sub = blk.ib % 8; + const uint is = sub * 2; + uint lo = uint(data_a_packed16[ib_k].scales[(is % 8) / 2]); + lo = (lo >> (4 * (is / 8))) & 0x0F0Fu; + uint hi = uint(data_a_packed16[ib_k].scales[(8 + (is % 4)) / 2]); + hi = (hi >> (2 * (is / 4))) & 0x0303u; + uint combined = lo | (hi << 4); + i8vec2 sc = unpack8(int32_t(combined)).xy; + float d = float(data_a_packed16[ib_k].d); + buf_a_d[(ks * KSCALES ) * BM + buf_ib] = d * float(int(sc.x) - 32); + buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = d * float(int(sc.y) - 32); + } +} + +#elif defined(DATA_A_NVFP4) + +struct block_a_prefetch { + uint32_t qs; + uint8_t d0; + uint8_t d1; +}; + +block_a_prefetch block_a_load(uint ib, uint loadr) { + block_a_prefetch blk; + const uint ib_k = ib / 2; + const uint ihalf = ib % 2; + const uint sub = ihalf * 2 + (loadr >> 1); + const uint byte_group = loadr & 1u; + + blk.qs = pack32(u8vec4(data_a[ib_k].qs[sub * 8 + byte_group * 4], + data_a[ib_k].qs[sub * 8 + byte_group * 4 + 1], + data_a[ib_k].qs[sub * 8 + byte_group * 4 + 2], + data_a[ib_k].qs[sub * 8 + byte_group * 4 + 3])); + blk.d0 = data_a[ib_k].d[ihalf * 2]; + blk.d1 = data_a[ib_k].d[ihalf * 2 + 1]; + + return blk; +} + +void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) { + const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F); + const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F); + const uint sub_base = (loadr >> 1) * 4; + const uint byte_group = loadr & 1u; + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + byte_group] = + pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y], + cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w])); + buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + 2 + byte_group] = + pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y], + cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w])); + + if (loadr == 0) { + buf_a_d[(ks * KSCALES ) * BM + buf_ib] = ue4m3_to_fp32(blk.d0) * 0.5; + buf_a_d[(ks * KSCALES + 1) * BM + buf_ib] = ue4m3_to_fp32(blk.d1) * 0.5; + } +} + +#endif + +// ===== B-side: load and store ===== + +struct block_b_prefetch { + ivec4 qs; + float16_t d; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + float16_t s; +#endif +}; + +block_b_prefetch block_b_load(uint ib_outer, uint ib_inner, uint loadr) { + block_b_prefetch blk; + blk.qs = data_b[ib_outer].qs[ib_inner * 2 + loadr]; + blk.d = data_b[ib_outer].ds[ib_inner].x; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + blk.s = data_b[ib_outer].ds[ib_inner].y; +#endif + return blk; +} + +void block_b_to_shmem(block_b_prefetch blk, uint buf_ib, uint ks, uint loadr, bool in_bounds) { + const ivec4 v = in_bounds ? blk.qs : ivec4(0); + const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr * 4; + buf_b_qs[base ] = v.x; + buf_b_qs[base + 1] = v.y; + buf_b_qs[base + 2] = v.z; + buf_b_qs[base + 3] = v.w; + if (loadr == 0) { + buf_b_d[ks * BN + buf_ib] = in_bounds ? float(blk.d) : 0.0f; +#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K) + buf_b_s[ks * BN + buf_ib] = in_bounds ? float(blk.s) : 0.0f; +#endif + } +} + +// ===== Framework macros ===== + +#ifdef MUL_MAT_ID +#define B_IB_CALC \ + const u16vec2 row_idx = row_ids[buf_ib]; \ + const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK \ + + (row_idx.x % p.ne11) * p.stride_b / BK; +#else +#define B_IB_CALC \ + const uint ib = pos_b_ib + buf_ib * p.stride_b / BK; +#endif + +#define PREFETCH_BLOCK(blk) \ + [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ + const uint buf_ib = loadc_a + li * loadstride_a; \ + if (buf_ib < BM) { \ + const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + pre_a[li * BK_STEP + ks] = block_a_load(ib + ks, loadr_a); \ + } \ + } \ + } \ + [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ + const uint buf_ib = loadc_b + li * loadstride_b; \ + if (buf_ib < BN) { \ + B_IB_CALC \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \ + pre_b[li * BK_STEP + ks] = block_b_load(ib_k / 4, ib_k % 4, loadr_b); \ + } \ + } \ + } + +#define STORE_BLOCK_TO_LDS(blk) \ + [[unroll]] for (uint li = 0; li < A_LOADS; li++) { \ + const uint buf_ib = loadc_a + li * loadstride_a; \ + if (buf_ib < BM) { \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + block_a_to_shmem(pre_a[li * BK_STEP + ks], buf_ib, ks, loadr_a); \ + } \ + } \ + } \ + [[unroll]] for (uint li = 0; li < B_LOADS; li++) { \ + const uint buf_ib = loadc_b + li * loadstride_b; \ + if (buf_ib < BN) { \ + [[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \ + const bool in_bounds = (blk) + ks * BK < end_k; \ + block_b_to_shmem(pre_b[li * BK_STEP + ks], buf_ib, ks, loadr_b, in_bounds); \ + } \ + } \ + } 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 d375c2d1277..27f56b3764c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -468,8 +468,9 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c base_dict["FLOAT16"] = "1"; } - base_dict["ACC_TYPE" ] = f16acc ? "float16_t" : "float"; - base_dict["ACC_TYPEV2"] = f16acc ? "f16vec2" : "vec2"; + base_dict["ACC_TYPE" ] = f16acc ? "float16_t" : "float"; + base_dict["ACC_TYPEV2" ] = f16acc ? "f16vec2" : "vec2"; + base_dict["ACC_TYPE_VEC4"] = f16acc ? "f16vec4" : "vec4"; if (f16acc) { base_dict["ACC_TYPE_MAX"] = "float16_t(65504.0)"; } @@ -627,6 +628,11 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"},}), fp16, coopmat, coopmat2, f16acc); } #endif + + if (coopmat && (tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1" || tname == "q8_0" || tname == "iq4_nl" || tname == "mxfp4" + || tname == "q3_k" || tname == "q4_k" || tname == "q5_k" || tname == "q6_k" || tname == "nvfp4")) { + string_to_spv(shader_name + "_" + tname + "_q8_1", "mul_mmq_cm1.comp", merge_maps(merge_maps(base_dict, float_type_dict), {{data_a_key, "1"}, {"D_TYPE", "float"}, {"D_TYPE_VEC4", "vec4"}}), fp16, coopmat, coopmat2, f16acc); + } } }