// Matmul-specific declarations and helpers
#ifdef VEC
#define VEC_SIZE 4u
#define SHMEM_TYPE vec4<f32>
#define DST_TYPE vec4<f32>
#ifdef INIT_SRC0_SHMEM_Q4_0
#define SRC0_TYPE u32
#else
#define SRC0_TYPE vec4<SRC0_INNER_TYPE>
#endif
#define SRC1_TYPE vec4<SRC1_INNER_TYPE>
fn store_shmem(val: vec4<f32>, idx: u32) {
shmem[idx] = val.x;
shmem[idx + 1] = val.y;
shmem[idx + 2] = val.z;
shmem[idx + 3] = val.w;
}
#else
#define VEC_SIZE 1u
#define SHMEM_TYPE f32
#define DST_TYPE f32
#define SRC0_TYPE SRC0_INNER_TYPE
#define SRC1_TYPE SRC1_INNER_TYPE
fn store_shmem(val: f32, idx: u32) {
shmem[idx] = val;
}
#endif
// Helpers for manual dequantization from u32-packed storage
#if defined(INIT_SRC0_SHMEM_Q4_0) || defined(INIT_SRC0_SHMEM_Q8_0) || defined(INIT_SRC0_SHMEM_Q6_K) || defined(INIT_SRC0_SHMEM_Q4_K)
fn load_src0_u16_at(byte_offset: u32) -> u32 {
let word = src0[byte_offset / 4u];
let shift = (byte_offset & 2u) * 8u;
return (word >> shift) & 0xFFFFu;
}
fn load_src0_u32_at(byte_offset: u32) -> u32 {
let word_idx = byte_offset / 4u;
let shift = (byte_offset & 3u) * 8u;
let lo = src0[word_idx];
if (shift == 0u) {
return lo;
}
let hi = src0[word_idx + 1u];
return (lo >> shift) | (hi << (32u - shift));
}
fn load_src0_f32_at(byte_offset: u32) -> f32 {
let packed = unpack2x16float(load_src0_u16_at(byte_offset));
return f32(packed[0]);
}
fn load_src0_byte_at(byte_offset: u32) -> u32 {
let word = src0[byte_offset / 4u];
return (word >> ((byte_offset & 3u) * 8u)) & 0xFFu;
}
#endif
#ifdef INIT_SRC0_SHMEM_FLOAT
fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
for (var elem_idx = thread_id * VEC_SIZE; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE * VEC_SIZE) {
let tile_m = elem_idx / TILE_SRC0_STRIDE;
let tile_k = elem_idx % TILE_SRC0_STRIDE;
let global_m = offset_m + tile_m;
let global_k = k_outer + tile_k;
let src0_idx = global_m * params.k + global_k;
let src0_val = select(
SRC0_TYPE(0.0),
src0[src0_idx/VEC_SIZE],
global_m < params.m && global_k < params.k);
store_shmem(SHMEM_TYPE(src0_val), elem_idx);
}
}
#endif
#ifdef INIT_SRC1_SHMEM_FLOAT
fn init_shmem_src1(thread_id: u32, offset_n: u32, k_outer: u32) {
for (var elem_idx = thread_id * VEC_SIZE; elem_idx < TILE_SRC1_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE * VEC_SIZE) {
let tile_n = elem_idx / TILE_K;
let tile_k = elem_idx % TILE_K;
let global_n = offset_n + tile_n;
let global_k = k_outer + tile_k;
let src1_idx = global_n * params.x_stride + global_k;
let src1_val = select(
SRC1_TYPE(0.0),
src1[src1_idx/VEC_SIZE],
global_n < params.n && global_k < params.k);
store_shmem(SHMEM_TYPE(src1_val), TILE_SRC0_SHMEM + elem_idx);
}
}
#endif
#ifdef INIT_SRC0_SHMEM_Q4_0
const BLOCK_SIZE = 32u;
const BLOCK_SIZE_BYTES = 18u;
fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
let BLOCKS_K = TILE_K / BLOCK_SIZE;
let PARAMS_BLOCKS_K = (params.k + BLOCK_SIZE - 1u) / BLOCK_SIZE;
let NQ = 16u;
let WEIGHTS_PER_F32 = 4u;
let F32_PER_THREAD = NQ / WEIGHTS_PER_F32;
for (var i = thread_id * NQ; i < TILE_SRC0_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F32;
let tile_m = blck_idx / BLOCKS_K;
let row_base = tile_m * TILE_SRC0_STRIDE;
let shmem_idx = row_base + block_offset * 2u;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
let global_k = k_outer / BLOCK_SIZE + block_k;
if (global_m < params.m && global_k < PARAMS_BLOCKS_K) {
let src0_idx = global_m * PARAMS_BLOCKS_K + global_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = load_src0_f32_at(block_byte_base);
for (var j = 0u; j < F32_PER_THREAD; j += 2u) {
let q_byte_offset = block_byte_base + 2u + 2u * (block_offset + j);
let q_packed = load_src0_u32_at(q_byte_offset);
for (var k = 0u; k < 4u; k++) {
let q_byte = get_byte(q_packed, k);
let q_hi = (f32((q_byte >> 4u) & 0xFu) - 8.0) * d;
let q_lo = (f32(q_byte & 0xFu) - 8.0) * d;
shmem[shmem_idx + j * 2u + k] = q_lo;
shmem[shmem_idx + j * 2u + k + 16u] = q_hi;
}
}
} else {
for (var j = 0u; j < 8u; j++) {
shmem[shmem_idx + j] = 0.0;
shmem[shmem_idx + j + 16u] = 0.0;
}
}
}
}
#endif
#ifdef INIT_SRC0_SHMEM_Q4_K
// Decode Q4_K weights straight into the f32 src0 shmem tile.
//
// Same motivation as the Q6_K loader below: `gemm_q4_k` is a batched-GEMV shape
// that re-dequantizes every weight once per token, so it buys submit count but no
// compute. Q4_K is the *bulk* of a Q4_K_M model (attn q/k/v/o, ffn gate/up, the
// conv projections), so leaving it on that shape kept the whole batched path
// slower than the per-token fallback even after Q6_K was reg-tiled.
//
// Q4_K super-block: 256 elems / 144 B — d f16 @0, dmin f16 @2,
// scales[12] (6-bit packed sub-scales + mins) @4, qs[128] @16.
// out[64j + l] = d*sc[2j] * (qs[32j+l] & 0xF) - dmin*mn[2j]
// out[64j + l + 32] = d*sc[2j+1] * (qs[32j+l] >> 4 ) - dmin*mn[2j+1]
// Inverting that per output index y: j = y/64, r = y%64, the low nibble and
// sub-block 2j for r < 32, else the high nibble and 2j+1.
const Q4K_BLOCK_SIZE = 256u;
const Q4K_BLOCK_BYTES = 144u;
// 6-bit sub-scale / min unpack — port of `decode_q4km_scales` (quant.rs).
fn q4k_sc(sb: u32, sub: u32) -> u32 {
if (sub < 4u) {
return load_src0_byte_at(sb + sub) & 63u;
}
return (load_src0_byte_at(sb + sub + 4u) & 0x0Fu)
| ((load_src0_byte_at(sb + sub - 4u) >> 6u) << 4u);
}
fn q4k_mn(sb: u32, sub: u32) -> u32 {
if (sub < 4u) {
return load_src0_byte_at(sb + sub + 4u) & 63u;
}
return (load_src0_byte_at(sb + sub + 4u) >> 4u)
| ((load_src0_byte_at(sb + sub) >> 6u) << 4u);
}
fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
let PARAMS_BLOCKS_K = (params.k + Q4K_BLOCK_SIZE - 1u) / Q4K_BLOCK_SIZE;
for (var i = thread_id; i < TILE_SRC0_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
let tile_m = i / TILE_K;
let tile_k = i % TILE_K;
let shmem_idx = tile_m * TILE_SRC0_STRIDE + tile_k;
let global_m = offset_m + tile_m;
let global_k = k_outer + tile_k;
if (global_m < params.m && global_k < params.k) {
let blk = global_k / Q4K_BLOCK_SIZE;
let y = global_k % Q4K_BLOCK_SIZE;
let base = (global_m * PARAMS_BLOCKS_K + blk) * Q4K_BLOCK_BYTES;
let d = load_src0_f32_at(base);
let dmin = load_src0_f32_at(base + 2u);
let sb = base + 4u;
let qs = base + 16u;
let j = y / 64u;
let r = y % 64u;
let l = r % 32u;
let hi = r >= 32u;
let sub = 2u * j + select(0u, 1u, hi);
let qb = load_src0_byte_at(qs + 32u * j + l);
let nib = select(qb & 0x0Fu, qb >> 4u, hi);
let sc = f32(q4k_sc(sb, sub));
let mn = f32(q4k_mn(sb, sub));
shmem[shmem_idx] = d * sc * f32(nib) - dmin * mn;
} else {
shmem[shmem_idx] = 0.0;
}
}
}
#endif
#ifdef INIT_SRC0_SHMEM_Q6_K
// Decode Q6_K weights straight into the f32 src0 shmem tile, so the
// dtype-agnostic register-tiled loop reads plain f32.
//
// This is the point of the whole kernel: the *batched-GEMV*-shaped gemm_q6_k
// re-dequantizes every weight once per token, so at n=512 the weights are decoded
// 512 times and it is measurably SLOWER than the per-token fallback it replaces
// (Mac prefill 42 vs 54 tok/s). Here each weight is decoded once per k-tile into
// shmem and then reused by all WORKGROUP_SIZE_N (=32) token columns.
//
// Q6_K super-block: 256 elems / 210 B — ql[128] @0, qh[64] @128,
// scales[16] (i8) @192, d f16 @208. The packing is *not* contiguous per output
// index, so we invert `dequantize_q6_k_block`'s (n, l, j) mapping per element:
// y = n*128 + j*32 + l with n in 0..2, j in 0..4, l in 0..32
// ql byte : n*64 + l (+32 when j is odd); low nibble for j<2, high for j>=2
// qh byte : n*32 + l; bits [2j, 2j+1]
// scale : n*8 + l/16 + 2j (signed i8)
// One output element per thread-iteration keeps every lane busy regardless of how
// TILE_K divides the 256-element super-block.
const Q6K_BLOCK_SIZE = 256u;
const Q6K_BLOCK_BYTES = 210u;
fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
let PARAMS_BLOCKS_K = (params.k + Q6K_BLOCK_SIZE - 1u) / Q6K_BLOCK_SIZE;
for (var i = thread_id; i < TILE_SRC0_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
let tile_m = i / TILE_K;
let tile_k = i % TILE_K;
let shmem_idx = tile_m * TILE_SRC0_STRIDE + tile_k;
let global_m = offset_m + tile_m;
let global_k = k_outer + tile_k;
if (global_m < params.m && global_k < params.k) {
let blk = global_k / Q6K_BLOCK_SIZE;
let y = global_k % Q6K_BLOCK_SIZE;
let base = (global_m * PARAMS_BLOCKS_K + blk) * Q6K_BLOCK_BYTES;
let n = y / 128u;
let r = y % 128u;
let j = r / 32u;
let l = r % 32u;
let ql_off = base + n * 64u + l + select(0u, 32u, (j & 1u) == 1u);
let qh_off = base + 128u + n * 32u + l;
let sc_off = base + 192u + n * 8u + (l / 16u) + 2u * j;
let qlb = load_src0_byte_at(ql_off);
let qhb = load_src0_byte_at(qh_off);
let nib = select(qlb & 0x0Fu, qlb >> 4u, j >= 2u);
let q = i32(nib | (((qhb >> (2u * j)) & 3u) << 4u)) - 32;
// Sub-scales are signed 8-bit; reading them unsigned flips every
// negative scale and yields plausible-but-wrong logits.
let sc_b = load_src0_byte_at(sc_off);
let sc = i32(sc_b << 24u) >> 24u;
let d = load_src0_f32_at(base + 208u);
shmem[shmem_idx] = d * f32(sc) * f32(q);
} else {
shmem[shmem_idx] = 0.0;
}
}
}
#endif
#ifdef INIT_SRC0_SHMEM_Q8_0
const Q8_BLOCK_SIZE = 32u;
const Q8_BLOCK_BYTES = 34u;
// Decode Q8_0 weight blocks (2-byte f16 scale + 32 int8 quants) straight into
// the f32 src0 shmem tile, so the dtype-agnostic register-tiled matmul loop
// reads plain f32. Work is split into 4-quant groups (8 per block) so every
// thread participates — one full block per thread would idle 7/8 of a 256-wide
// workgroup over a 32-row × 32-K tile.
const Q8_GROUP = 4u;
const Q8_GROUPS_PER_BLOCK = Q8_BLOCK_SIZE / Q8_GROUP;
fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
let BLOCKS_K = TILE_K / Q8_BLOCK_SIZE;
let PARAMS_BLOCKS_K = (params.k + Q8_BLOCK_SIZE - 1u) / Q8_BLOCK_SIZE;
let total_groups = TILE_SRC0_ROWS * BLOCKS_K * Q8_GROUPS_PER_BLOCK;
for (var g = thread_id; g < total_groups; g += TOTAL_WORKGROUP_SIZE) {
let blk = g / Q8_GROUPS_PER_BLOCK;
let grp = g % Q8_GROUPS_PER_BLOCK; // quants [grp*4, grp*4+4)
let tile_m = blk / BLOCKS_K;
let block_k = blk % BLOCKS_K;
let shmem_base = tile_m * TILE_SRC0_STRIDE + block_k * Q8_BLOCK_SIZE + grp * Q8_GROUP;
let global_m = offset_m + tile_m;
let global_k_block = k_outer / Q8_BLOCK_SIZE + block_k;
if (global_m < params.m && global_k_block < PARAMS_BLOCKS_K) {
let src0_block_idx = global_m * PARAMS_BLOCKS_K + global_k_block;
let block_byte_base = src0_block_idx * Q8_BLOCK_BYTES;
let d = load_src0_f32_at(block_byte_base);
let q_packed = load_src0_u32_at(block_byte_base + 2u + grp * Q8_GROUP);
for (var kk = 0u; kk < Q8_GROUP; kk++) {
let q_byte = get_byte(q_packed, kk);
// Sign-extend the int8 quant via arithmetic shift.
let q = i32(q_byte << 24u) >> 24u;
shmem[shmem_base + kk] = f32(q) * d;
}
} else {
for (var kk = 0u; kk < Q8_GROUP; kk++) {
shmem[shmem_base + kk] = 0.0;
}
}
}
}
#endif