llama-cpp-sys-4 0.7.0

Low Level Bindings to llama.cpp
Documentation
#version 450

#extension GL_EXT_control_flow_attributes : require
#extension GL_KHR_shader_subgroup_basic : require
#extension GL_KHR_shader_subgroup_shuffle : require

// 16 lanes per token, indexed idst + hc*isrc: idst in bits 0..1, isrc in bits 2..3,
// so subgroupShuffleXor by 1|2 reduces a row and by 4|8 a column.

layout(constant_id = 0) const uint SUBGROUP_SIZE = 32;

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

layout(push_constant) uniform parameter
{
    uint n_tokens;

    uint nbm0; uint nbm1;   // mixes
    uint nbs0;              // scale
    uint nbb0;              // base
    uint nbd0; uint nbd1; uint nbd2;   // dst

    uint m_offset;
    uint s_offset;
    uint b_offset;
    uint d_offset;

    float eps;
    uint n_iter;
};

layout(binding = 0, std430) readonly buffer M { float data_m[]; };
layout(binding = 1, std430) readonly buffer S { float data_s[]; };
layout(binding = 2, std430) readonly buffer B { float data_b[]; };
layout(binding = 3, std430) writeonly buffer D { float data_d[]; };

const uint hc          = 4;
const uint comb_offset = 2 * hc;

const uint TOKENS_PER_SUBGROUP = SUBGROUP_SIZE / 16;

void main() {
    const uint lane = gl_SubgroupInvocationID;
    const uint blk  = lane >> 4;    // which 16-lane block, i.e. which token
    const uint idx  = lane & 15;    // idst + hc*isrc

    const uint sg = gl_WorkGroupID.x * gl_WorkGroupSize.y + gl_SubgroupID;
    const uint it = sg * TOKENS_PER_SUBGROUP + blk;

    // no early return, the shuffles need every lane; out-of-range blocks compute a discarded value
    const bool in_range = it < n_tokens;

    const float scale_comb = data_s[s_offset + 2 * nbs0];

    float v = 0.0f;
    if (in_range) {
        v = data_m[m_offset + (comb_offset + idx) * nbm0 + it * nbm1] * scale_comb
          + data_b[b_offset + (comb_offset + idx) * nbb0];
    }

    // Softmax across destinations: the four lanes sharing an isrc.
    float vmax = max(v, subgroupShuffleXor(v, 1));
    vmax = max(vmax, subgroupShuffleXor(vmax, 2));
    v = exp(v - vmax);

    float sum = v + subgroupShuffleXor(v, 1);
    sum += subgroupShuffleXor(sum, 2);
    v = v / sum + eps;

    // Normalize columns: equal destination indices are four lanes apart.
    sum = v + subgroupShuffleXor(v, 4);
    sum += subgroupShuffleXor(sum, 8);
    v /= sum + eps;

    for (uint i = 1; i < n_iter; ++i) {
        sum = v + subgroupShuffleXor(v, 1);
        sum += subgroupShuffleXor(sum, 2);
        v /= sum + eps;

        sum = v + subgroupShuffleXor(v, 4);
        sum += subgroupShuffleXor(sum, 8);
        v /= sum + eps;
    }

    if (in_range) {
        const uint idst = idx & 3;
        const uint isrc = idx >> 2;
        data_d[d_offset + idst * nbd0 + isrc * nbd1 + it * nbd2] = v;
    }
}