Compare commits

...
Author SHA1 Message Date
Ruben Ortlam 965e57103f fix shmem support function, clean up comments 2026-08-29 14:05:48 +02:00
Ruben Ortlam 8b194a15eb adapt to upstream changes 2026-08-29 11:51:15 +02:00
Ruben Ortlam e85ff84334 use BK_STEP 2 on MUL_MAT_ID 2026-08-29 11:51:15 +02:00
Ruben Ortlam 584ae05d5f rdna4 tuning 2026-08-29 11:51:15 +02:00
Ruben Ortlam 04d7ef7b24 fix iq4_nl and nvfp4 performance 2026-08-29 11:51:15 +02:00
Ruben Ortlam 6421fbe532 clean up 2026-08-29 11:51:15 +02:00
Ruben Ortlam c30a8740f2 improve offset application 2026-08-29 11:51:15 +02:00
Ruben Ortlam 0eab519531 add RDNA4 architecture, use for hardcoded coopmat elem thread access, set BK_STEP back to 4 2026-08-29 11:51:15 +02:00
Ruben Ortlam f252045d89 undo uint8_t, gate to RDNA3/4 2026-08-29 11:51:15 +02:00
Ruben Ortlam 9ef5996da4 merge shmem arrays 2026-08-29 11:51:15 +02:00
Ruben Ortlam 2ba1d225e1 dedup b scales 2026-08-29 11:51:15 +02:00
Ruben Ortlam eeeacedd4e improvements 2026-08-29 11:51:15 +02:00
Ruben Ortlam 092661df98 improve performance 2026-08-29 11:51:15 +02:00
Ruben Ortlam 3985f64670 improve performance 2026-08-29 11:51:15 +02:00
Ruben Ortlam 22f588c10d fix l warptile 2026-08-29 11:51:15 +02:00
Ruben Ortlam 4695f108bb add q3_k, q4_k, q5_k, q6_k and nvfp4 support 2026-08-29 11:51:15 +02:00
Ruben Ortlam ab94ffdd30 use 4-byte loads where possible 2026-08-29 11:51:15 +02:00
Ruben Ortlam 5037f027aa use shmem arrays for LUTs 2026-08-29 11:51:15 +02:00
Ruben Ortlam ba8402e9f6 remove elem row/col fast path, invalid for RDNA4 2026-08-29 11:51:15 +02:00
Ruben Ortlam dc08fa9101 support iq4_nl and mxfp4 2026-08-29 11:51:15 +02:00
Ruben Ortlam 3cd76432d7 fix mul_mat_id bug 2026-08-29 11:51:15 +02:00
Ruben Ortlam d8ed37209a fix segfault 2026-08-29 11:51:15 +02:00
Ruben Ortlam fa37db93ce enable mul_mat_id support 2026-08-29 11:51:15 +02:00
Ruben Ortlam ed294e8724 restructure mmq cm1 functions 2026-08-29 11:51:15 +02:00
Ruben Ortlam d2a5ac69c9 add q4_1, q5_0, q5_1 support 2026-08-29 11:51:15 +02:00
Ruben Ortlam eb46b8cf0a move quant-specific prefetch function out of main file 2026-08-29 11:51:15 +02:00
Ruben Ortlam 1825270eb4 fix compilation 2026-08-29 11:51:15 +02:00
Ruben Ortlam 42a6269bda use BK_STEP 4 2026-08-29 11:51:15 +02:00
Ruben Ortlam ab6e84bbdd only force subgroup size 32 on AMD RDNA 2026-08-29 11:51:15 +02:00
Ruben Ortlam 57dd3f8bd3 skip computation for inactive tiles 2026-08-29 11:51:15 +02:00
Ruben Ortlam 2b9c53381e restructure for vgpr use 2026-08-29 11:51:15 +02:00
Ruben Ortlam ef5c62992f Revert "increase large tile size"
This reverts commit 7fabc25c5e.
2026-08-29 11:51:15 +02:00
Ruben Ortlam a9e0e89c0b increase large tile size 2026-08-29 11:51:15 +02:00
Ruben Ortlam 57b4276784 use wave32 2026-08-29 11:51:15 +02:00
Ruben Ortlam 15b3ce0de3 Revert "revert load reordering and scale pre-loading"
This reverts commit fbaefe0eaa.
2026-08-29 11:51:15 +02:00
Ruben Ortlam 9f069f42e8 clean up 2026-08-29 11:51:15 +02:00
Ruben Ortlam 23356762ca workgroup scheduling for cache proximity 2026-08-29 11:51:15 +02:00
Ruben Ortlam 31b9bc683e revert load reordering and scale pre-loading 2026-08-29 11:51:15 +02:00
Ruben Ortlam ad015d6a44 add faster RDNA int->float conversion 2026-08-29 11:51:15 +02:00
Ruben Ortlam 95d14c99ea use float for scales 2026-08-29 11:51:15 +02:00
Ruben Ortlam 14a6080dd5 coopmat load first, then wmma 2026-08-29 11:51:15 +02:00
Ruben Ortlam b453be38a3 preload scales 2026-08-29 11:51:15 +02:00
Ruben OrtlamandPiotr Wilkin 5553b89912 double buffering
Co-authored-by: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>
2026-08-29 11:51:13 +02:00
Ruben OrtlamandPiotr Wilkin 2a878d9ea9 use larger workgroups
Co-authored-by: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>
2026-08-29 11:51:10 +02:00
Ruben OrtlamandPiotr Wilkin b121ef3b33 add BK_STEP to shader, default to 2
Co-authored-by: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>
2026-08-29 11:51:08 +02:00
Ruben Ortlam 513f43600c add q8_0 support 2026-08-29 11:51:08 +02:00
Ruben OrtlamandPiotr Wilkin 162560885a probe and directly access coopmat values instead of going through shmem
Co-authored-by: Piotr Wilkin (ilintar) <piotr.wilkin@syndatis.com>
2026-08-29 11:51:04 +02:00
Ruben Ortlam 5f034d6086 use scalar sums 2026-08-29 11:30:17 +02:00
Ruben Ortlam e086b9ca05 apply scales inline 2026-08-29 11:30:17 +02:00
Ruben Ortlam 13f6d8e389 vulkan: add int8 coopmat quantized matmul shader 2026-08-29 11:30:17 +02:00
4 changed files with 1268 additions and 23 deletions
+196 -21
View File
@@ -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<uint32_t>& 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<uint32_t, 3> 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;
const bool cm1_k_tile = (t == GGML_TYPE_Q3_K || t == GGML_TYPE_Q6_K ||
t == GGML_TYPE_NVFP4);
if (!ggml_vk_matmul_int_shmem_support(device, s_int, false, t)) {
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<uint32_t>& 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:
@@ -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<int8_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> 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<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> 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<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
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<int8_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> cache_a[cms_per_row * K_SUB];
coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> 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<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
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<int8_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> cache_a[cms_per_row * K_SUB];
coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> 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<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
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
}
@@ -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); \
} \
} \
}
@@ -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);
}
}
}