mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-07 16:37:57 +02:00
Compare commits
50
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
965e57103f | ||
|
|
8b194a15eb | ||
|
|
e85ff84334 | ||
|
|
584ae05d5f | ||
|
|
04d7ef7b24 | ||
|
|
6421fbe532 | ||
|
|
c30a8740f2 | ||
|
|
0eab519531 | ||
|
|
f252045d89 | ||
|
|
9ef5996da4 | ||
|
|
2ba1d225e1 | ||
|
|
eeeacedd4e | ||
|
|
092661df98 | ||
|
|
3985f64670 | ||
|
|
22f588c10d | ||
|
|
4695f108bb | ||
|
|
ab94ffdd30 | ||
|
|
5037f027aa | ||
|
|
ba8402e9f6 | ||
|
|
dc08fa9101 | ||
|
|
3cd76432d7 | ||
|
|
d8ed37209a | ||
|
|
fa37db93ce | ||
|
|
ed294e8724 | ||
|
|
d2a5ac69c9 | ||
|
|
eb46b8cf0a | ||
|
|
1825270eb4 | ||
|
|
42a6269bda | ||
|
|
ab6e84bbdd | ||
|
|
57dd3f8bd3 | ||
|
|
2b9c53381e | ||
|
|
ef5c62992f | ||
|
|
a9e0e89c0b | ||
|
|
57b4276784 | ||
|
|
15b3ce0de3 | ||
|
|
9f069f42e8 | ||
|
|
23356762ca | ||
|
|
31b9bc683e | ||
|
|
ad015d6a44 | ||
|
|
95d14c99ea | ||
|
|
14a6080dd5 | ||
|
|
b453be38a3 | ||
|
|
5553b89912 | ||
|
|
2a878d9ea9 | ||
|
|
b121ef3b33 | ||
|
|
513f43600c | ||
|
|
162560885a | ||
|
|
5f034d6086 | ||
|
|
e086b9ca05 | ||
|
|
13f6d8e389 |
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user