// 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
#ifdef INIT_SRC0_SHMEM_Q4_0
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]);
}
#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