cera 0.1.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
#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