llama-cpp-sys-4 0.7.0

Low Level Bindings to llama.cpp
Documentation
#version 450

#extension GL_EXT_control_flow_attributes : require

// Fan one stream back out to hc streams and add the combination-weighted
// residuals:
//
//   dst[i0, idst, it] = x[i0, it]*post[idst, it]
//                     + sum_isrc residual[i0, isrc, it]*comb[idst, isrc, it]

layout(constant_id = 0) const uint BLOCK_SIZE = 256;

layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;

layout(push_constant) uniform parameter
{
    uint n_embd;
    uint n_tokens;

    uint nbx0; uint nbx1;              // x
    uint nbr0; uint nbr1; uint nbr2;   // residual
    uint nbp0; uint nbp1;              // post
    uint nbc0; uint nbc1; uint nbc2;   // comb
    uint nbd0; uint nbd1; uint nbd2;   // dst

    uint x_offset;
    uint r_offset;
    uint p_offset;
    uint c_offset;
    uint d_offset;
};

layout(binding = 0, std430) readonly buffer X { float data_x[]; };
layout(binding = 1, std430) readonly buffer R { float data_r[]; };
layout(binding = 2, std430) readonly buffer P { float data_p[]; };
layout(binding = 3, std430) readonly buffer C { float data_c[]; };
layout(binding = 4, std430) writeonly buffer D { float data_d[]; };

const uint hc = 4;

shared float post_s[hc];
shared float comb_s[hc * hc];

void main() {
    const uint tid = gl_LocalInvocationID.x;
    const uint it  = gl_WorkGroupID.y;

    if (tid < hc) {
        post_s[tid] = data_p[p_offset + tid * nbp0 + it * nbp1];
    }
    if (tid < hc * hc) {
        const uint idst = tid & 3;
        const uint isrc = tid >> 2;
        comb_s[tid] = data_c[c_offset + idst * nbc0 + isrc * nbc1 + it * nbc2];
    }
    barrier();

    // After the barrier, so every invocation reaches it.
    const uint i0 = gl_WorkGroupID.x * BLOCK_SIZE + tid;
    if (i0 >= n_embd) {
        return;
    }

    const float xv = data_x[x_offset + i0 * nbx0 + it * nbx1];

    const uint rb = r_offset + i0 * nbr0 + it * nbr2;

    float r[hc];
    [[unroll]]
    for (uint isrc = 0; isrc < hc; ++isrc) {
        r[isrc] = data_r[rb + isrc * nbr1];
    }

    [[unroll]]
    for (uint idst = 0; idst < hc; ++idst) {
        float result = xv * post_s[idst];
        [[unroll]]
        for (uint isrc = 0; isrc < hc; ++isrc) {
            result = fma(r[isrc], comb_s[idst + hc * isrc], result);
        }
        data_d[d_offset + i0 * nbd0 + idst * nbd1 + it * nbd2] = result;
    }
}