llama-cpp-sys-4 0.7.0

Low Level Bindings to llama.cpp
Documentation
#version 450
#extension GL_EXT_shader_explicit_arithmetic_types : require

#include "mul_mat_vec_base.glsl"

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

FLOAT_TYPE temp[NUM_COLS][NUM_ROWS];

// Walks the packed bytes directly (byte m, digit t) rather than via
// tq1_0_byte_of()/tq1_0_digit_of(): one byte per thread, expanded in place.
void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
    uint a_offset, b_offset, d_offset;
    get_offsets(a_offset, b_offset, d_offset);

    const uint num_blocks_per_row = p.ncols / QUANT_K;
    const uint tid = gl_LocalInvocationID.x;

    [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
        [[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) {
            temp[j][i] = FLOAT_TYPE(0);
        }
    }

    for (uint nrow = 0; nrow < num_rows; ++nrow) {
        const uint ib0 = a_offset + (first_row + nrow) * num_blocks_per_row;
        for (uint jcol = 0; jcol < NUM_COLS; ++jcol) {
            const uint b_base = (jcol * p.batch_stride_b);
            for (uint i = tid/8; i < num_blocks_per_row; i += gl_WorkGroupSize.x/8) {
                const FLOAT_TYPE d = float(data_a[ib0 + i].d);

                // First qs chunk: 32 bytes (5*32 elements)
                [[unroll]] for (uint m = tid%8; m < 32; m += 8) {
                    const uint q_byte = uint(data_a[ib0 + i].qs[m]);
                    [[unroll]] for (uint t = 0; t < 5; ++t) {
                        const uint xi = tq1_0_trit(q_byte, t);
                        const FLOAT_TYPE dequant_val = FLOAT_TYPE(d * (float(xi) - 1.0f));
                        const uint elem = t * 32u + m;
                        const uint b_idx = i * QUANT_K + elem;
                        temp[jcol][nrow] += dequant_val * FLOAT_TYPE(data_b[b_base + b_offset + b_idx]);
                    }
                }

                // Second qs chunk: 16 bytes (5*16 elements)
                [[unroll]] for (uint m = tid%8; m < 16; m += 8) {
                    const uint q_byte = uint(data_a[ib0 + i].qs[32u + m]);
                    [[unroll]] for (uint t = 0; t < 5; ++t) {
                        const uint xi = tq1_0_trit(q_byte, t);
                        const FLOAT_TYPE dequant_val = FLOAT_TYPE(d * (float(xi) - 1.0f));
                        const uint elem = 160u + t * 16u + m;
                        const uint b_idx = i * QUANT_K + elem;
                        temp[jcol][nrow] += dequant_val * FLOAT_TYPE(data_b[b_base + b_offset + b_idx]);
                    }
                }

                // qh bytes: 4 bytes (4*4 elements)
                [[unroll]] for (uint j = tid%8; j < 4; j += 8) {
                    const uint qh_byte = uint(data_a[ib0 + i].qh[j]);
                    [[unroll]] for (uint t = 0; t < 4; ++t) {
                        const uint xi = tq1_0_trit(qh_byte, t);
                        const FLOAT_TYPE dequant_val = FLOAT_TYPE(d * (float(xi) - 1.0f));
                        const uint elem = 240u + t * 4u + j;
                        const uint b_idx = i * QUANT_K + elem;
                        temp[jcol][nrow] += dequant_val * FLOAT_TYPE(data_b[b_base + b_offset + b_idx]);
                    }
                }
            }
        }
    }

    reduce_result(temp, d_offset, first_row, num_rows, tid);
}

void main() {
    const uint first_row = NUM_ROWS * (gl_WorkGroupID.x + gl_NumWorkGroups.x * gl_WorkGroupID.z);

    if (first_row + NUM_ROWS <= p.stride_d) {
        compute_outputs(first_row, NUM_ROWS);
    } else {
        if (first_row >= p.stride_d) {
            return;
        }
        compute_outputs(first_row, p.stride_d - first_row);
    }
}