cera 0.5.3

Rust-native LLM inference engine
Documentation
// slang-entries: gemv_q4_0_qkv
//
// Fused 3-in-1 Q/K/V Q4_0 Matrix-Vector Multiply (GEMV) in Slang for both Metal and WGSL.
// Computes Q (M_q x K), K (M_kv x K), and V (M_kv x K) projections in a single dispatch pass,
// reusing cooperatively staged activation inputs in workgroup shared memory.
// Tiled across NR=4 output rows per workgroup.

[[vk::binding(0)]] StructuredBuffer<uint>    w_q    : register(t0);
[[vk::binding(1)]] StructuredBuffer<uint>    w_k    : register(t1);
[[vk::binding(2)]] StructuredBuffer<uint>    w_v    : register(t2);
[[vk::binding(3)]] StructuredBuffer<float>   x      : register(t3);
[[vk::binding(4)]] RWStructuredBuffer<float> y_q    : register(u4);
[[vk::binding(5)]] RWStructuredBuffer<float> y_k    : register(u5);
[[vk::binding(6)]] RWStructuredBuffer<float> y_v    : register(u6);
[[vk::binding(7)]] StructuredBuffer<uint4>   params : register(t7);

static const uint NR = 4u;
static const uint NQ = 16u;
static const uint WG_SIZE = 32u;
static const uint BLOCK_BYTES = 18u;
static const uint CHUNK_BLOCKS = 16u;
static const uint CHUNK_FLOATS = CHUNK_BLOCKS * 32u; // 512 floats (2 KB)

groupshared float partials[NR * WG_SIZE];
groupshared float x_stage[CHUNK_FLOATS];

uint get_wid(uint3 wid) {
    return wid.x + wid.y * 65535u;
}

float block_scale_q(uint blk_byte) {
    uint word = w_q[blk_byte >> 2u];
    uint scale_bits = ((blk_byte & 2u) != 0u) ? (word >> 16u) : (word & 0xFFFFu);
    return float(f16tof32(scale_bits));
}

float block_scale_k(uint blk_byte) {
    uint word = w_k[blk_byte >> 2u];
    uint scale_bits = ((blk_byte & 2u) != 0u) ? (word >> 16u) : (word & 0xFFFFu);
    return float(f16tof32(scale_bits));
}

float block_scale_v(uint blk_byte) {
    uint word = w_v[blk_byte >> 2u];
    uint scale_bits = ((blk_byte & 2u) != 0u) ? (word >> 16u) : (word & 0xFFFFu);
    return float(f16tof32(scale_bits));
}

uint q_pair_q(uint qs_byte) {
    uint word = w_q[qs_byte >> 2u];
    return ((qs_byte & 2u) != 0u) ? (word >> 16u) : (word & 0xFFFFu);
}

uint q_pair_k(uint qs_byte) {
    uint word = w_k[qs_byte >> 2u];
    return ((qs_byte & 2u) != 0u) ? (word >> 16u) : (word & 0xFFFFu);
}

uint q_pair_v(uint qs_byte) {
    uint word = w_v[qs_byte >> 2u];
    return ((qs_byte & 2u) != 0u) ? (word >> 16u) : (word & 0xFFFFu);
}

[shader("compute")]
[numthreads(WG_SIZE, 1, 1)]
void gemv_q4_0_qkv(
    uint3 lid : SV_GroupThreadID,
    uint3 wid : SV_GroupID
) {
    uint m_q = params[0].x;
    uint m_kv = params[0].y;
    uint k = params[0].z;
    uint nb = k / 32u;
    uint row_bytes = nb * BLOCK_BYTES;
    uint global_r0 = get_wid(wid) * NR;
    uint tid = lid.x;

    uint ix = tid / 2u;
    uint il = (tid & 1u) * 8u;

    uint which_matrix; // 0 = Q, 1 = K, 2 = V
    uint r0;
    uint m_cur;
    if (global_r0 < m_q) {
        which_matrix = 0u;
        r0 = global_r0;
        m_cur = m_q;
    } else if (global_r0 < m_q + m_kv) {
        which_matrix = 1u;
        r0 = global_r0 - m_q;
        m_cur = m_kv;
    } else {
        which_matrix = 2u;
        r0 = global_r0 - m_q - m_kv;
        m_cur = m_kv;
    }

    float sumf[NR];
    [unroll]
    for (uint r = 0u; r < NR; r++) {
        sumf[r] = 0.0f;
    }

    for (uint chunk_b = 0u; chunk_b < nb; chunk_b += CHUNK_BLOCKS) {
        uint chunk_k_start = chunk_b * 32u;
        uint chunk_k_len = min(CHUNK_FLOATS, k - chunk_k_start);
        uint chunk_k_vec4 = chunk_k_len / 4u;
        for (uint i = tid; i < chunk_k_vec4; i += WG_SIZE) {
            uint base_idx = chunk_k_start + i * 4u;
            x_stage[i * 4u + 0u] = x[base_idx + 0u];
            x_stage[i * 4u + 1u] = x[base_idx + 1u];
            x_stage[i * 4u + 2u] = x[base_idx + 2u];
            x_stage[i * 4u + 3u] = x[base_idx + 3u];
        }
        GroupMemoryBarrierWithGroupSync();

        uint ib_local = ix;
        while (ib_local < CHUNK_BLOCKS && (chunk_b + ib_local) < nb) {
            uint ib = chunk_b + ib_local;
            uint yb_stage_off = ib_local * 32u + il;

            float a0  = x_stage[yb_stage_off + 0u];
            float a1  = x_stage[yb_stage_off + 1u];
            float a2  = x_stage[yb_stage_off + 2u];
            float a3  = x_stage[yb_stage_off + 3u];
            float a4  = x_stage[yb_stage_off + 4u];
            float a5  = x_stage[yb_stage_off + 5u];
            float a6  = x_stage[yb_stage_off + 6u];
            float a7  = x_stage[yb_stage_off + 7u];
            float a8  = x_stage[yb_stage_off + 16u];
            float a9  = x_stage[yb_stage_off + 17u];
            float a10 = x_stage[yb_stage_off + 18u];
            float a11 = x_stage[yb_stage_off + 19u];
            float a12 = x_stage[yb_stage_off + 20u];
            float a13 = x_stage[yb_stage_off + 21u];
            float a14 = x_stage[yb_stage_off + 22u];
            float a15 = x_stage[yb_stage_off + 23u];

            float y0 = a0;
            float y1 = a1 / 256.0f;
            float y2 = a2;
            float y3 = a3 / 256.0f;
            float y4 = a4;
            float y5 = a5 / 256.0f;
            float y6 = a6;
            float y7 = a7 / 256.0f;
            float y8 = a8 / 16.0f;
            float y9 = a9 / 4096.0f;
            float y10 = a10 / 16.0f;
            float y11 = a11 / 4096.0f;
            float y12 = a12 / 16.0f;
            float y13 = a13 / 4096.0f;
            float y14 = a14 / 16.0f;
            float y15 = a15 / 4096.0f;

            float sumy0 = (a0 + a1) + (a2 + a3) + (a4 + a5) + (a6 + a7);
            float sumy1 = (a8 + a9) + (a10 + a11) + (a12 + a13) + (a14 + a15);
            float sumy_total = sumy0 + sumy1;

            if (which_matrix == 0u) {
                [unroll]
                for (uint r = 0u; r < NR; r++) {
                    if (r0 + r >= m_cur) { continue; }
                    uint blk_byte = (r0 + r) * row_bytes + ib * BLOCK_BYTES;
                    float d = block_scale_q(blk_byte);
                    uint qs_byte = blk_byte + 2u + il;

                    float acc0 = 0.0f;
                    float acc1 = 0.0f;
                    float acc2 = 0.0f;
                    float acc3 = 0.0f;

                    uint q0 = q_pair_q(qs_byte + 0u);
                    acc0 += y0 * float(q0 & 0x000Fu);
                    acc1 += y1 * float(q0 & 0x0F00u);
                    acc2 += y8 * float(q0 & 0x00F0u);
                    acc3 += y9 * float(q0 & 0xF000u);

                    uint q1 = q_pair_q(qs_byte + 2u);
                    acc0 += y2 * float(q1 & 0x000Fu);
                    acc1 += y3 * float(q1 & 0x0F00u);
                    acc2 += y10 * float(q1 & 0x00F0u);
                    acc3 += y11 * float(q1 & 0xF000u);

                    uint q2 = q_pair_q(qs_byte + 4u);
                    acc0 += y4 * float(q2 & 0x000Fu);
                    acc1 += y5 * float(q2 & 0x0F00u);
                    acc2 += y12 * float(q2 & 0x00F0u);
                    acc3 += y13 * float(q2 & 0xF000u);

                    uint q3 = q_pair_q(qs_byte + 6u);
                    acc0 += y6 * float(q3 & 0x000Fu);
                    acc1 += y7 * float(q3 & 0x0F00u);
                    acc2 += y14 * float(q3 & 0x00F0u);
                    acc3 += y15 * float(q3 & 0xF000u);

                    sumf[r] += d * (sumy_total * -8.0f + acc0 + acc1 + acc2 + acc3);
                }
            } else if (which_matrix == 1u) {
                [unroll]
                for (uint r = 0u; r < NR; r++) {
                    if (r0 + r >= m_cur) { continue; }
                    uint blk_byte = (r0 + r) * row_bytes + ib * BLOCK_BYTES;
                    float d = block_scale_k(blk_byte);
                    uint qs_byte = blk_byte + 2u + il;

                    float acc0 = 0.0f;
                    float acc1 = 0.0f;
                    float acc2 = 0.0f;
                    float acc3 = 0.0f;

                    uint q0 = q_pair_k(qs_byte + 0u);
                    acc0 += y0 * float(q0 & 0x000Fu);
                    acc1 += y1 * float(q0 & 0x0F00u);
                    acc2 += y8 * float(q0 & 0x00F0u);
                    acc3 += y9 * float(q0 & 0xF000u);

                    uint q1 = q_pair_k(qs_byte + 2u);
                    acc0 += y2 * float(q1 & 0x000Fu);
                    acc1 += y3 * float(q1 & 0x0F00u);
                    acc2 += y10 * float(q1 & 0x00F0u);
                    acc3 += y11 * float(q1 & 0xF000u);

                    uint q2 = q_pair_k(qs_byte + 4u);
                    acc0 += y4 * float(q2 & 0x000Fu);
                    acc1 += y5 * float(q2 & 0x0F00u);
                    acc2 += y12 * float(q2 & 0x00F0u);
                    acc3 += y13 * float(q2 & 0xF000u);

                    uint q3 = q_pair_k(qs_byte + 6u);
                    acc0 += y6 * float(q3 & 0x000Fu);
                    acc1 += y7 * float(q3 & 0x0F00u);
                    acc2 += y14 * float(q3 & 0x00F0u);
                    acc3 += y15 * float(q3 & 0xF000u);

                    sumf[r] += d * (sumy_total * -8.0f + acc0 + acc1 + acc2 + acc3);
                }
            } else {
                [unroll]
                for (uint r = 0u; r < NR; r++) {
                    if (r0 + r >= m_cur) { continue; }
                    uint blk_byte = (r0 + r) * row_bytes + ib * BLOCK_BYTES;
                    float d = block_scale_v(blk_byte);
                    uint qs_byte = blk_byte + 2u + il;

                    float acc0 = 0.0f;
                    float acc1 = 0.0f;
                    float acc2 = 0.0f;
                    float acc3 = 0.0f;

                    uint q0 = q_pair_v(qs_byte + 0u);
                    acc0 += y0 * float(q0 & 0x000Fu);
                    acc1 += y1 * float(q0 & 0x0F00u);
                    acc2 += y8 * float(q0 & 0x00F0u);
                    acc3 += y9 * float(q0 & 0xF000u);

                    uint q1 = q_pair_v(qs_byte + 2u);
                    acc0 += y2 * float(q1 & 0x000Fu);
                    acc1 += y3 * float(q1 & 0x0F00u);
                    acc2 += y10 * float(q1 & 0x00F0u);
                    acc3 += y11 * float(q1 & 0xF000u);

                    uint q2 = q_pair_v(qs_byte + 4u);
                    acc0 += y4 * float(q2 & 0x000Fu);
                    acc1 += y5 * float(q2 & 0x0F00u);
                    acc2 += y12 * float(q2 & 0x00F0u);
                    acc3 += y13 * float(q2 & 0xF000u);

                    uint q3 = q_pair_v(qs_byte + 6u);
                    acc0 += y6 * float(q3 & 0x000Fu);
                    acc1 += y7 * float(q3 & 0x0F00u);
                    acc2 += y14 * float(q3 & 0x00F0u);
                    acc3 += y15 * float(q3 & 0xF000u);

                    sumf[r] += d * (sumy_total * -8.0f + acc0 + acc1 + acc2 + acc3);
                }
            }

            ib_local += NQ;
        }
        GroupMemoryBarrierWithGroupSync();
    }

    [unroll]
    for (uint r = 0u; r < NR; r++) {
        partials[r * WG_SIZE + tid] = sumf[r];
    }
    GroupMemoryBarrierWithGroupSync();

    for (uint stride = WG_SIZE / 2u; stride > 0u; stride >>= 1) {
        if (tid < stride) {
            [unroll]
            for (uint r = 0u; r < NR; r++) {
                uint idx = r * WG_SIZE + tid;
                partials[idx] += partials[idx + stride];
            }
        }
        GroupMemoryBarrierWithGroupSync();
    }

    if (tid == 0u) {
        [unroll]
        for (uint r = 0u; r < NR; r++) {
            if (r0 + r < m_cur) {
                if (which_matrix == 0u) {
                    y_q[r0 + r] = partials[r * WG_SIZE];
                } else if (which_matrix == 1u) {
                    y_k[r0 + r] = partials[r * WG_SIZE];
                } else {
                    y_v[r0 + r] = partials[r * WG_SIZE];
                }
            }
        }
    }
}