cera 0.3.0

Rust-native LLM inference engine
Documentation
// 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)
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

#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