#version 450
// Subgroup-reduced Q5_K matrix-vector product: y[n] = sum_k W[n,k]*x[k], W stored as GGUF Q5_K
// super-blocks kept verbatim in VRAM (256 weights / 176 bytes = 44 u32 per super-block: u32[0]=
// {d,dmin} f16; u32[1..4]=12 scale bytes; u32[4..12]=32 qh high-bit bytes; u32[12..44]=128 qs
// low-nibble bytes). Decode is memory-bandwidth bound on this ~256 GB/s APU. This variant fuses the
// per-row reduction into one subgroup arithmetic reduce (subgroupAdd) instead of one thread per row,
// so a whole subgroup streams one row's super-blocks in parallel — more memory-level parallelism per
// row, no shared-memory barrier. Dispatched ONLY when the device advertises subgroup ARITHMETIC
// support (gated host-side); the scalar mul_mat_vec_q5k kernel is the fallback. Decode is
// bit-identical to mul_mat_vec_q5k (CPU k_quants BlockQ5K::to_float).
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
#extension GL_KHR_shader_subgroup_basic : require
#extension GL_KHR_shader_subgroup_arithmetic : require
layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in;
layout(set = 0, binding = 0) readonly buffer W { uint w[]; }; // Q5_K blocks, 44 u32 / 256-block
layout(set = 0, binding = 1) readonly buffer X { float x[]; }; // activation vector, length k
layout(set = 0, binding = 2) writeonly buffer Y { float y[]; }; // output, length nout
// woff: u32 offset into w[] where this weight matrix starts (0 for a plain 2D weight; e*n*(k/256)*44
// to select expert e of a resident MoE bank). k is a multiple of 256.
layout(push_constant) uniform Pc { uint nout; uint k; uint woff; };
const uint QK_K = 256u;
const uint BLK_U32 = 44u; // 176 bytes / 4
const uint QH_U32 = 4u; // u32 index where the 32 qh bytes start
const uint QS_U32 = 12u; // u32 index where the 128 qs bytes start
uint scale_byte(uint base, uint b) {
uint word = w[base + 1u + (b >> 2u)];
return (word >> ((b & 3u) * 8u)) & 0xFFu;
}
uint get_scale_min_k4(uint base, uint j) {
uint sc, m;
if (j < 4u) {
sc = scale_byte(base, j) & 63u;
m = scale_byte(base, j + 4u) & 63u;
} else {
uint s_j4 = scale_byte(base, j + 4u);
uint s_jm4 = scale_byte(base, j - 4u);
uint s_j = scale_byte(base, j);
sc = (s_j4 & 0x0Fu) | ((s_jm4 >> 6u) << 4u);
m = (s_j4 >> 4u) | ((s_j >> 6u) << 4u);
}
return (sc << 8u) | m;
}
uint qh_byte(uint base, uint b) {
uint word = w[base + QH_U32 + (b >> 2u)];
return (word >> ((b & 3u) * 8u)) & 0xFFu;
}
uint qs_byte(uint base, uint b) {
uint word = w[base + QS_U32 + (b >> 2u)];
return (word >> ((b & 3u) * 8u)) & 0xFFu;
}
void main() {
// One subgroup per output row. Row = global subgroup index.
uint n = gl_WorkGroupID.x * gl_NumSubgroups + gl_SubgroupID;
if (n >= nout) {
return;
}
uint nblocks = k / QK_K;
uint rowbase = woff + n * nblocks * BLK_U32; // u32 offset of this row within the selected matrix
uint lane = gl_SubgroupInvocationID;
uint lanes = gl_SubgroupSize;
float acc = 0.0;
for (uint blk = lane; blk < nblocks; blk += lanes) {
uint base = rowbase + blk * BLK_U32;
float d = float(unpackHalf2x16(w[base]).x);
float dmin = float(unpackHalf2x16(w[base]).y);
uint xblk = blk * QK_K; // activation offset for this super-block
for (uint ci = 0u; ci < 4u; ci++) {
uint is = ci * 2u;
uint sm1 = get_scale_min_k4(base, is);
uint sm2 = get_scale_min_k4(base, is + 1u);
float d1 = d * float(sm1 >> 8u);
float m1 = dmin * float(sm1 & 0xFFu);
float d2 = d * float(sm2 >> 8u);
float m2 = dmin * float(sm2 & 0xFFu);
uint qoff = ci * 32u; // byte offset into qs for this chunk
uint u1 = 1u << (ci * 2u); // high-bit mask for the low nibble
uint u2 = 2u << (ci * 2u); // high-bit mask for the high nibble
uint xlo = xblk + ci * 64u; // lower-nibble outputs land here
uint xhi = xlo + 32u; // upper-nibble outputs
for (uint l = 0u; l < 32u; l++) {
uint qval = qs_byte(base, qoff + l);
uint hbit = qh_byte(base, l);
float addlo = ((hbit & u1) != 0u) ? 16.0 : 0.0;
float addhi = ((hbit & u2) != 0u) ? 16.0 : 0.0;
float wlo = d1 * (float(qval & 0x0Fu) + addlo) - m1;
float whi = d2 * (float(qval >> 4u) + addhi) - m2;
acc += wlo * x[xlo + l];
acc += whi * x[xhi + l];
}
}
}
float total = subgroupAdd(acc);
if (subgroupElect()) {
y[n] = total;
}
}