cera 0.5.1

Rust-native LLM inference engine
Documentation
// src0 (weight) staging for `mul_mat_reg_tile.wgsl`.
//
// Every loader decodes weights to f32 into the k-major `sa` tile, so the
// register-tiled inner loop is dtype-agnostic and reads plain f32. Each weight
// is decoded once per k-tile and then reused by all TILE_COLS token columns —
// that reuse is the entire reason this kernel beats the batched-GEMV-shaped
// `gemm_*` kernels, which re-dequantize per token.
//
// The layout is `sa[k][m]` (stride SA_STRIDE), NOT `sa[m][k]`. Staging writes
// therefore transpose. That is deliberate: the write happens once per element
// per k-tile while the inner loop reads each element TILE_COLS times, and the
// read side is what wants the TILE_M rows of one k to be consecutive.
//
// Requires from the includer: `sa`, `SA_STRIDE`, `TILE_ROWS`, `TILE_K`,
// `TOTAL_WORKGROUP_SIZE`, `params`, `src0`, and `get_byte`.

/// Transposing store into the k-major src0 tile.
fn store_sa(tile_m: u32, tile_k: u32, value: f32) {
    sa[tile_k * SA_STRIDE + tile_m] = value;
}

// 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) || defined(INIT_SRC0_SHMEM_Q5_K)
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 word = src0[byte_offset / 4u];
    let half = (word >> ((byte_offset & 2u) * 8u)) & 0xFFFFu;
    return unpack2x16float(half)[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 i = thread_id; i < TILE_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
        let tile_m = i / TILE_K;
        let tile_k = i % TILE_K;
        let global_m = offset_m + tile_m;
        let global_k = k_outer + tile_k;
        store_sa(tile_m, tile_k, select(
            0.0,
            src0[global_m * params.k + global_k],
            global_m < params.m && global_k < params.k,
        ));
    }
}
#endif

#ifdef INIT_SRC0_SHMEM_Q4_0
// Q4_0 super-block: 32 elems / 18 B — d f16 @0, then 16 bytes where byte j
// holds weight j in the low nibble and weight j+16 in the high nibble.
//
// Blocked rather than per-element: each thread stages 8 consecutive k of one
// row, which is exactly one aligned nibble-half of the packed bytes, so the f16
// scale and the two u32 quant words are read once per 8 outputs instead of once
// each.
const Q4_0_BLOCK_SIZE = 32u;
const Q4_0_BLOCK_BYTES = 18u;
const Q4_0_PER_THREAD = 8u;

fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
    let blocks_k = (params.k + Q4_0_BLOCK_SIZE - 1u) / Q4_0_BLOCK_SIZE;

    for (var i = thread_id * Q4_0_PER_THREAD;
         i < TILE_ROWS * TILE_K;
         i += TOTAL_WORKGROUP_SIZE * Q4_0_PER_THREAD) {
        let tile_m = i / TILE_K;
        let tile_k = i % TILE_K;
        let global_m = offset_m + tile_m;
        let global_k = k_outer + tile_k;

        // Position of `global_k` within its 32-element block. `i` steps by 8 and
        // TILE_K is a multiple of 8 (checked host-side), so `w` is a multiple of
        // 8 and the 8 weights are one nibble-half of bytes [w % 16, w % 16 + 8),
        // never straddling a block.
        let w = global_k % Q4_0_BLOCK_SIZE;
        let base = (global_m * blocks_k + global_k / Q4_0_BLOCK_SIZE) * Q4_0_BLOCK_BYTES;
        let d = select(0.0, load_src0_f32_at(base), global_m < params.m && global_k < params.k);
        let q_lo = load_src0_u32_at(base + 2u + (w % 16u));
        let q_hi = load_src0_u32_at(base + 6u + (w % 16u));

        for (var j = 0u; j < Q4_0_PER_THREAD; j++) {
            // `j & 3u` is `j` for j < 4 and `j - 4` for j >= 4 — the byte index
            // within whichever u32 the select picks. Spelled this way rather
            // than `j - 4u` so the unused arm never evaluates an out-of-range
            // shift (WGSL leaves shifts >= 32 indeterminate).
            let byte = select(get_byte(q_lo, j & 3u), get_byte(q_hi, j & 3u), j >= 4u);
            let nib = select(byte & 0xFu, byte >> 4u, w >= 16u);
            // Zero past the edge per element, not per 8-run: `d` is gated on the
            // run's FIRST k, so with a `params.k` that is not a multiple of 8 the
            // trailing lanes would otherwise stage decoded garbage. The k loop
            // runs a full TILE_K with no tail check and relies on this.
            let live = global_m < params.m && global_k + j < params.k;
            store_sa(tile_m, tile_k + j, select(0.0, (f32(nib) - 8.0) * d, live));
        }
    }
}
#endif

#ifdef INIT_SRC0_SHMEM_Q8_0
// Q8_0 super-block: 32 elems / 34 B — d f16 @0, then 32 int8 quants.
const Q8_0_BLOCK_SIZE = 32u;
const Q8_0_BLOCK_BYTES = 34u;

fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
    let blocks_k = (params.k + Q8_0_BLOCK_SIZE - 1u) / Q8_0_BLOCK_SIZE;

    for (var i = thread_id; i < TILE_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
        let tile_m = i / TILE_K;
        let tile_k = i % TILE_K;
        let global_m = offset_m + tile_m;
        let global_k = k_outer + tile_k;

        var v = 0.0;
        if (global_m < params.m && global_k < params.k) {
            let base = (global_m * blocks_k + global_k / Q8_0_BLOCK_SIZE) * Q8_0_BLOCK_BYTES;
            // Quants are signed 8-bit; sign-extend via arithmetic shift.
            let q = i32(load_src0_byte_at(base + 2u + global_k % Q8_0_BLOCK_SIZE) << 24u) >> 24u;
            v = f32(q) * load_src0_f32_at(base);
        }
        store_sa(tile_m, tile_k, v);
    }
}
#endif

#ifdef INIT_SRC0_SHMEM_Q4_K
// 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_base: u32, sub: u32) -> u32 {
    if (sub < 4u) {
        return load_src0_byte_at(sb_base + sub) & 63u;
    }
    return (load_src0_byte_at(sb_base + sub + 4u) & 0x0Fu)
        | ((load_src0_byte_at(sb_base + sub - 4u) >> 6u) << 4u);
}

fn q4k_mn(sb_base: u32, sub: u32) -> u32 {
    if (sub < 4u) {
        return load_src0_byte_at(sb_base + sub + 4u) & 63u;
    }
    return (load_src0_byte_at(sb_base + sub + 4u) >> 4u)
        | ((load_src0_byte_at(sb_base + sub) >> 6u) << 4u);
}

fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
    let blocks_k = (params.k + Q4K_BLOCK_SIZE - 1u) / Q4K_BLOCK_SIZE;

    for (var i = thread_id; i < TILE_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
        let tile_m = i / TILE_K;
        let tile_k = i % TILE_K;
        let global_m = offset_m + tile_m;
        let global_k = k_outer + tile_k;

        var v = 0.0;
        if (global_m < params.m && global_k < params.k) {
            let base = (global_m * blocks_k + global_k / Q4K_BLOCK_SIZE) * Q4K_BLOCK_BYTES;
            let y = global_k % Q4K_BLOCK_SIZE;
            let j = y / 64u;
            let r = y % 64u;
            let hi = r >= 32u;
            let sub = 2u * j + select(0u, 1u, hi);
            let sb_base = base + 4u;

            let qb = load_src0_byte_at(base + 16u + 32u * j + r % 32u);
            let nib = select(qb & 0x0Fu, qb >> 4u, hi);

            v = load_src0_f32_at(base) * f32(q4k_sc(sb_base, sub)) * f32(nib)
                - load_src0_f32_at(base + 2u) * f32(q4k_mn(sb_base, sub));
        }
        store_sa(tile_m, tile_k, v);
    }
}
#endif

#ifdef INIT_SRC0_SHMEM_Q6_K
// 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)
const Q6K_BLOCK_SIZE = 256u;
const Q6K_BLOCK_BYTES = 210u;

fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
    let blocks_k = (params.k + Q6K_BLOCK_SIZE - 1u) / Q6K_BLOCK_SIZE;

    for (var i = thread_id; i < TILE_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
        let tile_m = i / TILE_K;
        let tile_k = i % TILE_K;
        let global_m = offset_m + tile_m;
        let global_k = k_outer + tile_k;

        var v = 0.0;
        if (global_m < params.m && global_k < params.k) {
            let base = (global_m * blocks_k + global_k / Q6K_BLOCK_SIZE) * Q6K_BLOCK_BYTES;
            let y = global_k % Q6K_BLOCK_SIZE;
            let n = y / 128u;
            let r = y % 128u;
            let j = r / 32u;
            let l = r % 32u;

            let qlb = load_src0_byte_at(base + n * 64u + l + select(0u, 32u, (j & 1u) == 1u));
            let qhb = load_src0_byte_at(base + 128u + n * 32u + l);
            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(base + 192u + n * 8u + (l / 16u) + 2u * j);
            let sc = i32(sc_b << 24u) >> 24u;

            v = load_src0_f32_at(base + 208u) * f32(sc) * f32(q);
        }
        store_sa(tile_m, tile_k, v);
    }
}
#endif

#ifdef INIT_SRC0_SHMEM_Q5_K
// Q5_K super-block: 256 elems / 176 B — d f16 @0, dmin f16 @2,
// scales[12] (6-bit packed sub-scales + mins, identical to Q4_K) @4,
// qh[32] (the 5th-bit plane) @16, qs[128] (low 4 bits) @48.
//   out[64j + l]      = d*sc[2j]   * ((qs[32j+l] & 0xF) + fifth) - dmin*mn[2j]
//   out[64j + l + 32] = d*sc[2j+1] * ((qs[32j+l] >> 4 ) + fifth) - dmin*mn[2j+1]
// where `fifth` is 16 when the qh bit for that element is set, else 0 (the code
// local `hi` below is a different thing: the bool high-nibble selector r>=32).
// The qh plane is indexed by `l` (= y%32) alone — the same 32 bytes serve all `j`
// sub-blocks, consuming bit `sub` (= 2j for the low nibble, 2j+1 for the high).
// So this is the Q4_K loader with `qs` at byte 48 instead of 16 and the extra
// `+fifth` term; the 6-bit scale/min unpack is byte-for-byte the same.
// Inverting per output index y: j = y/64, r = y%64, low nibble and sub-block 2j
// for r < 32, else the high nibble and 2j+1; l = y%32.
const Q5K_BLOCK_SIZE = 256u;
const Q5K_BLOCK_BYTES = 176u;

// 6-bit sub-scale / min unpack — identical to `q4k_sc`/`q4k_mn`, redeclared here
// because each pipeline compiles the template with exactly one loader define, so
// the Q4_K helpers are not in scope for the Q5_K build.
fn q5k_sc(sb_base: u32, sub: u32) -> u32 {
    if (sub < 4u) {
        return load_src0_byte_at(sb_base + sub) & 63u;
    }
    return (load_src0_byte_at(sb_base + sub + 4u) & 0x0Fu)
        | ((load_src0_byte_at(sb_base + sub - 4u) >> 6u) << 4u);
}

fn q5k_mn(sb_base: u32, sub: u32) -> u32 {
    if (sub < 4u) {
        return load_src0_byte_at(sb_base + sub + 4u) & 63u;
    }
    return (load_src0_byte_at(sb_base + sub + 4u) >> 4u)
        | ((load_src0_byte_at(sb_base + sub) >> 6u) << 4u);
}

fn init_shmem_src0(thread_id: u32, offset_m: u32, k_outer: u32) {
    let blocks_k = (params.k + Q5K_BLOCK_SIZE - 1u) / Q5K_BLOCK_SIZE;

    for (var i = thread_id; i < TILE_ROWS * TILE_K; i += TOTAL_WORKGROUP_SIZE) {
        let tile_m = i / TILE_K;
        let tile_k = i % TILE_K;
        let global_m = offset_m + tile_m;
        let global_k = k_outer + tile_k;

        var v = 0.0;
        if (global_m < params.m && global_k < params.k) {
            let base = (global_m * blocks_k + global_k / Q5K_BLOCK_SIZE) * Q5K_BLOCK_BYTES;
            let y = global_k % Q5K_BLOCK_SIZE;
            let j = y / 64u;
            let r = y % 64u;
            let hi = r >= 32u;
            let sub = 2u * j + select(0u, 1u, hi);
            let l = y % 32u;
            let sb_base = base + 4u;

            let qb = load_src0_byte_at(base + 48u + 32u * j + l);
            let nib = select(qb & 0x0Fu, qb >> 4u, hi);
            // 5th bit: qh byte at base+16+l, bit `sub`; contributes +16 when set.
            let qhb = load_src0_byte_at(base + 16u + l);
            let fifth = f32(((qhb >> sub) & 1u) * 16u);

            v = load_src0_f32_at(base) * f32(q5k_sc(sb_base, sub)) * (f32(nib) + fifth)
                - load_src0_f32_at(base + 2u) * f32(q5k_mn(sb_base, sub));
        }
        store_sa(tile_m, tile_k, v);
    }
}
#endif