#version 450
// Q4_K matrix-matrix product (the prefill/compute path): y[m,n] = sum_k x[m,k]*W[n,k] with W stored
// as GGUF Q4_K super-blocks kept verbatim in VRAM, SAME layout and decode as mul_mat_vec_q4k (256
// weights / 144 bytes = 36 u32 per super-block). Prefill has M>1 activation rows sharing one weight
// matrix, so the old path dequantized the whole weight to f32 every forward (massive VRAM + bandwidth)
// This kernel keeps W quantized and decodes each weight value ONCE per output column, reusing it
// across all M rows: weight traffic equals a single matvec while M rows are produced. One invocation
// computes one output column n for up to MAX_M rows (the host tiles M by MAX_M). Decode MUST match the
// CPU k_quants BlockQ4K::to_float (identical to mul_mat_vec_q4k).
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : 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[]; }; // Q4_K blocks, 36 u32 / 256-block
layout(set = 0, binding = 1) readonly buffer X { float x[]; }; // activations, row-major [M, k]
layout(set = 0, binding = 2) writeonly buffer Y { float y[]; }; // output, row-major [M, nout]
// m0: first activation row this dispatch handles; mcount: rows in this tile (1..MAX_M); nout: total
// columns; k: contraction dim (multiple of 256). x row stride is k, y row stride is nout. woff: u32
// offset into w[] where this weight matrix starts (0 for a plain 2D weight; e*nout*(k/256)*36 to
// select expert e of a resident MoE bank).
layout(push_constant) uniform Pc { uint m0; uint mcount; uint nout; uint k; uint woff; };
const uint QK_K = 256u;
const uint BLK_U32 = 36u; // 144 bytes / 4
const uint MAX_M = 8u; // register accumulators per invocation; host tiles M by this
// Byte `b` (0-based) out of the block's scale region (u32[1..4], i.e. the 12 scale bytes).
uint scale_byte(uint base, uint b) {
uint word = w[base + 1u + (b >> 2u)];
return (word >> ((b & 3u) * 8u)) & 0xFFu;
}
// get_scale_min_k4 from k_quants/utils.rs: returns packed (sc<<8 | m), each 6-bit.
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;
}
void main() {
uint n = gl_GlobalInvocationID.x;
if (n >= nout) {
return;
}
uint nblocks = k / QK_K;
uint rowbase = woff + n * nblocks * BLK_U32; // u32 offset of weight row n within the selected matrix
float acc[MAX_M];
for (uint r = 0u; r < MAX_M; r++) {
acc[r] = 0.0;
}
for (uint blk = 0u; blk < nblocks; blk++) {
uint base = rowbase + blk * BLK_U32;
float d = float(unpackHalf2x16(w[base]).x);
float dmin = float(unpackHalf2x16(w[base]).y);
uint qbase = base + 4u; // u32 index where the 128 quant bytes start
uint xblk = blk * QK_K; // activation offset for this super-block
// 4 chunks of 64 weights; chunk c uses scale-pair (2c, 2c+1) and qs bytes [c*32, c*32+32).
for (uint c = 0u; c < 4u; c++) {
uint is = c * 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 = c * 32u; // byte offset into qs for this chunk
uint xlo = xblk + c * 64u; // lower-nibble outputs land here
uint xhi = xlo + 32u; // upper-nibble outputs
for (uint l = 0u; l < 32u; l++) {
uint qb = qoff + l;
uint qword = w[qbase + (qb >> 2u)];
uint q = (qword >> ((qb & 3u) * 8u)) & 0xFFu;
// Decode the two weights of this byte ONCE, then reuse across all rows.
float wlo = d1 * float(q & 0x0Fu) - m1;
float whi = d2 * float(q >> 4u) - m2;
for (uint r = 0u; r < mcount; r++) {
uint xrow = (m0 + r) * k;
acc[r] += wlo * x[xrow + xlo + l];
acc[r] += whi * x[xrow + xhi + l];
}
}
}
}
for (uint r = 0u; r < mcount; r++) {
y[(m0 + r) * nout + n] = acc[r];
}
}