hanzo-ml 0.11.85

Fast multi-backend tensor & ML framework for Rust (CPU/CUDA/Metal/Vulkan/ROCm) with quantization — the compute core of the Hanzo stack.
Documentation
#version 450
// IQ2_XXS matrix-vector product (decode): y[n] = sum_k W[n,k]*x[k]. GGUF IQ2_XXS super-block packs
// 256 weights into 66 bytes: d (f16, byte 0..2) + qs[32] (u16, byte 2..66). 66 is not u32-aligned,
// so the loader repacks each block to a padded 68-byte / 17-u32 stride (trailing 2 bytes zero). Per
// 32-weight sub-block e: aux0 = u32 at byte 2+8e, aux1 = u32 at byte 2+8e+4; the 4-bit scale =
// aux1>>28 (db = d*(0.5+scale)*0.25). Lane (0..31): group l=lane>>3 selects grid index gi =
// (aux0>>8l)&0xFF and a sign byte KSIGNS[(aux1>>7l)&127]; element j=lane&7 reads grid point
// g = byte j of the 8-byte grid entry, sign = bit j of the sign byte. weight = db*g*(+/-1). This is
// byte-exact with k_quants/iq_quants BlockIQ2xxs::to_float. One invocation computes one output row.
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in;

layout(set = 0, binding = 0) readonly buffer W { uint   w[]; };  // IQ2_XXS blocks, 17 u32 / 256-block
layout(set = 0, binding = 1) readonly buffer X { float  x[]; };  // activation vector, length k
layout(set = 0, binding = 2) writeonly buffer Y { float y[]; };  // output, length nout
layout(push_constant) uniform Pc { uint nout; uint k; };         // k is a multiple of 256

const uint BLK_U32 = 17u;   // 68 bytes / 4 (66 real + 2 pad)

// IQ2XXS_GRID: 256 u64 entries (2 u32 each); KSIGNS: 128 bytes
const uint IQ2XXS_GRID[512] = uint[512](
    0x08080808u, 0x08080808u, 0x0808082bu, 0x08080808u, 0x08081919u, 0x08080808u, 0x08082b08u, 0x08080808u, 0x08082b2bu,
    0x08080808u, 0x08190819u, 0x08080808u, 0x08191908u, 0x08080808u, 0x082b0808u, 0x08080808u, 0x082b082bu, 0x08080808u,
    0x082b2b08u, 0x08080808u, 0x082b2b2bu, 0x08080808u, 0x19080819u, 0x08080808u, 0x19081908u, 0x08080808u, 0x19190808u,
    0x08080808u, 0x19192b08u, 0x08080808u, 0x192b0819u, 0x08080808u, 0x192b1908u, 0x08080808u, 0x2b080808u, 0x08080808u,
    0x2b08082bu, 0x08080808u, 0x2b082b2bu, 0x08080808u, 0x2b2b082bu, 0x08080808u, 0x08080819u, 0x08080819u, 0x08081908u,
    0x08080819u, 0x08190808u, 0x08080819u, 0x08191919u, 0x08080819u, 0x19080808u, 0x08080819u, 0x2b081908u, 0x08080819u,
    0x2b192b08u, 0x08080819u, 0x08080808u, 0x0808082bu, 0x0808082bu, 0x0808082bu, 0x082b082bu, 0x0808082bu, 0x2b08082bu,
    0x0808082bu, 0x08080819u, 0x08081908u, 0x08081908u, 0x08081908u, 0x08190808u, 0x08081908u, 0x082b0819u, 0x08081908u,
    0x082b1908u, 0x08081908u, 0x19080808u, 0x08081908u, 0x1908082bu, 0x08081908u, 0x19082b08u, 0x08081908u, 0x192b0808u,
    0x08081908u, 0x2b080819u, 0x08081908u, 0x2b081908u, 0x08081908u, 0x2b190808u, 0x08081908u, 0x2b2b1908u, 0x08081908u,
    0x08080808u, 0x08081919u, 0x0808082bu, 0x08081919u, 0x08082b08u, 0x08081919u, 0x082b0808u, 0x08081919u, 0x1908192bu,
    0x08081919u, 0x192b2b19u, 0x08081919u, 0x2b080808u, 0x08081919u, 0x2b190819u, 0x08081919u, 0x08082b19u, 0x0808192bu,
    0x08190808u, 0x0808192bu, 0x19080808u, 0x0808192bu, 0x2b081908u, 0x0808192bu, 0x2b2b1908u, 0x0808192bu, 0x08080808u,
    0x08082b08u, 0x08081919u, 0x08082b08u, 0x08082b08u, 0x08082b08u, 0x08191908u, 0x08082b08u, 0x082b2b08u, 0x08082b08u,
    0x19080819u, 0x08082b08u, 0x19081908u, 0x08082b08u, 0x19190808u, 0x08082b08u, 0x1919082bu, 0x08082b08u, 0x2b082b08u,
    0x08082b08u, 0x08081908u, 0x08082b19u, 0x19080808u, 0x08082b19u, 0x0808082bu, 0x08082b2bu, 0x08191908u, 0x08082b2bu,
    0x08080819u, 0x08190808u, 0x08081908u, 0x08190808u, 0x08190808u, 0x08190808u, 0x082b0819u, 0x08190808u, 0x19080808u,
    0x08190808u, 0x192b0808u, 0x08190808u, 0x2b081908u, 0x08190808u, 0x2b190808u, 0x08190808u, 0x2b191919u, 0x08190808u,
    0x08080808u, 0x08190819u, 0x08082b08u, 0x08190819u, 0x082b0808u, 0x08190819u, 0x19190808u, 0x08190819u, 0x19192b2bu,
    0x08190819u, 0x2b080808u, 0x08190819u, 0x082b1908u, 0x0819082bu, 0x19081919u, 0x0819082bu, 0x08080808u, 0x08191908u,
    0x08082b08u, 0x08191908u, 0x082b0808u, 0x08191908u, 0x082b1919u, 0x08191908u, 0x19082b19u, 0x08191908u, 0x2b080808u,
    0x08191908u, 0x08192b08u, 0x08191919u, 0x192b082bu, 0x08191919u, 0x08080808u, 0x0819192bu, 0x0819192bu, 0x0819192bu,
    0x08080819u, 0x08192b08u, 0x08081908u, 0x08192b08u, 0x08190808u, 0x08192b08u, 0x19080808u, 0x08192b08u, 0x2b080819u,
    0x08192b08u, 0x08080808u, 0x08192b19u, 0x08081919u, 0x08192b19u, 0x2b2b0808u, 0x08192b19u, 0x19190819u, 0x08192b2bu,
    0x08080808u, 0x082b0808u, 0x0808082bu, 0x082b0808u, 0x08082b2bu, 0x082b0808u, 0x19081908u, 0x082b0808u, 0x192b0819u,
    0x082b0808u, 0x2b080808u, 0x082b0808u, 0x2b08082bu, 0x082b0808u, 0x082b2b19u, 0x082b0819u, 0x19082b08u, 0x082b0819u,
    0x08080808u, 0x082b082bu, 0x0808082bu, 0x082b082bu, 0x08080819u, 0x082b1908u, 0x08081908u, 0x082b1908u, 0x08190808u,
    0x082b1908u, 0x19080808u, 0x082b1908u, 0x1919192bu, 0x082b1908u, 0x08080808u, 0x082b1919u, 0x19080819u, 0x082b1919u,
    0x192b1908u, 0x082b1919u, 0x2b190808u, 0x082b192bu, 0x08082b08u, 0x082b2b08u, 0x082b0808u, 0x082b2b08u, 0x2b191908u,
    0x082b2b08u, 0x19081908u, 0x082b2b2bu, 0x08080819u, 0x19080808u, 0x08081908u, 0x19080808u, 0x08190808u, 0x19080808u,
    0x08192b08u, 0x19080808u, 0x082b0819u, 0x19080808u, 0x082b1908u, 0x19080808u, 0x19080808u, 0x19080808u, 0x19082b08u,
    0x19080808u, 0x1919192bu, 0x19080808u, 0x192b0808u, 0x19080808u, 0x2b080819u, 0x19080808u, 0x2b081908u, 0x19080808u,
    0x2b190808u, 0x19080808u, 0x08080808u, 0x19080819u, 0x082b0808u, 0x19080819u, 0x192b0819u, 0x19080819u, 0x2b080808u,
    0x19080819u, 0x2b081919u, 0x19080819u, 0x08080819u, 0x1908082bu, 0x08190808u, 0x1908082bu, 0x19082b08u, 0x1908082bu,
    0x1919192bu, 0x1908082bu, 0x192b2b08u, 0x1908082bu, 0x08080808u, 0x19081908u, 0x08082b08u, 0x19081908u, 0x082b0808u,
    0x19081908u, 0x2b080808u, 0x19081908u, 0x2b192b19u, 0x19081908u, 0x0819082bu, 0x19081919u, 0x082b1908u, 0x19081919u,
    0x08080808u, 0x1908192bu, 0x08080819u, 0x19082b08u, 0x08081908u, 0x19082b08u, 0x08190808u, 0x19082b08u, 0x19080808u,
    0x19082b08u, 0x19081919u, 0x19082b08u, 0x08080808u, 0x19082b19u, 0x19192b08u, 0x19082b19u, 0x192b0819u, 0x19082b19u,
    0x2b08082bu, 0x19082b19u, 0x19081919u, 0x19082b2bu, 0x2b190808u, 0x19082b2bu, 0x08080808u, 0x19190808u, 0x08082b08u,
    0x19190808u, 0x08190819u, 0x19190808u, 0x08192b19u, 0x19190808u, 0x082b0808u, 0x19190808u, 0x2b080808u, 0x19190808u,
    0x2b082b08u, 0x19190808u, 0x08081908u, 0x19190819u, 0x1908082bu, 0x19190819u, 0x2b2b1908u, 0x19190819u, 0x2b190819u,
    0x1919082bu, 0x2b190808u, 0x19191908u, 0x2b19082bu, 0x19191908u, 0x08082b2bu, 0x19191919u, 0x08080819u, 0x1919192bu,
    0x19191908u, 0x1919192bu, 0x08080808u, 0x19192b08u, 0x08190819u, 0x19192b08u, 0x08192b19u, 0x19192b08u, 0x192b1908u,
    0x19192b08u, 0x19080808u, 0x19192b19u, 0x08082b08u, 0x19192b2bu, 0x08081908u, 0x192b0808u, 0x08190808u, 0x192b0808u,
    0x19080808u, 0x192b0808u, 0x192b2b08u, 0x192b0808u, 0x08080808u, 0x192b0819u, 0x19191919u, 0x192b0819u, 0x08192b08u,
    0x192b082bu, 0x192b0808u, 0x192b082bu, 0x08080808u, 0x192b1908u, 0x08081919u, 0x192b1908u, 0x08190808u, 0x192b1919u,
    0x0819082bu, 0x192b1919u, 0x2b081908u, 0x192b1919u, 0x1908082bu, 0x192b2b08u, 0x08080808u, 0x2b080808u, 0x0808082bu,
    0x2b080808u, 0x08082b2bu, 0x2b080808u, 0x19080819u, 0x2b080808u, 0x2b08082bu, 0x2b080808u, 0x08081908u, 0x2b080819u,
    0x08192b08u, 0x2b080819u, 0x19080808u, 0x2b080819u, 0x08190819u, 0x2b08082bu, 0x08080819u, 0x2b081908u, 0x08081908u,
    0x2b081908u, 0x08190808u, 0x2b081908u, 0x08191919u, 0x2b081908u, 0x19080808u, 0x2b081908u, 0x192b0808u, 0x2b081908u,
    0x08080808u, 0x2b081919u, 0x1908192bu, 0x2b081919u, 0x2b191908u, 0x2b081919u, 0x08082b19u, 0x2b08192bu, 0x19080808u,
    0x2b08192bu, 0x192b0808u, 0x2b08192bu, 0x0808082bu, 0x2b082b08u, 0x08081908u, 0x2b082b19u, 0x08190819u, 0x2b082b2bu,
    0x08081908u, 0x2b190808u, 0x08190808u, 0x2b190808u, 0x082b1908u, 0x2b190808u, 0x19080808u, 0x2b190808u, 0x2b2b0819u,
    0x2b190808u, 0x0819192bu, 0x2b190819u, 0x2b080808u, 0x2b190819u, 0x19081919u, 0x2b19082bu, 0x08080808u, 0x2b191908u,
    0x082b082bu, 0x2b191908u, 0x19081908u, 0x2b191908u, 0x19190819u, 0x2b191919u, 0x2b080819u, 0x2b192b08u, 0x082b0808u,
    0x2b192b19u, 0x0808082bu, 0x2b2b0808u, 0x19190808u, 0x2b2b0808u, 0x2b081919u, 0x2b2b0808u, 0x08082b19u, 0x2b2b0819u,
    0x08080808u, 0x2b2b082bu, 0x08192b08u, 0x2b2b1908u, 0x19190808u, 0x2b2b2b08u, 0x08081908u, 0x2b2b2b19u
);

const uint KSIGNS[128] = uint[128](
    0u, 129u, 130u, 3u, 132u, 5u, 6u, 135u, 136u, 9u, 10u, 139u, 12u, 141u, 142u, 15u, 144u, 17u, 18u, 147u,
    20u, 149u, 150u, 23u, 24u, 153u, 154u, 27u, 156u, 29u, 30u, 159u, 160u, 33u, 34u, 163u, 36u, 165u, 166u,
    39u, 40u, 169u, 170u, 43u, 172u, 45u, 46u, 175u, 48u, 177u, 178u, 51u, 180u, 53u, 54u, 183u, 184u, 57u, 58u,
    187u, 60u, 189u, 190u, 63u, 192u, 65u, 66u, 195u, 68u, 197u, 198u, 71u, 72u, 201u, 202u, 75u, 204u, 77u,
    78u, 207u, 80u, 209u, 210u, 83u, 212u, 85u, 86u, 215u, 216u, 89u, 90u, 219u, 92u, 221u, 222u, 95u, 96u, 225u,
    226u, 99u, 228u, 101u, 102u, 231u, 232u, 105u, 106u, 235u, 108u, 237u, 238u, 111u, 240u, 113u, 114u, 243u,
    116u, 245u, 246u, 119u, 120u, 249u, 250u, 123u, 252u, 125u, 126u, 255u
);


// Unsigned byte `b` (0-based) within the block at u32 index `base`.
uint byte_u(uint base, uint b) {
    return (w[base + (b >> 2u)] >> ((b & 3u) * 8u)) & 0xFFu;
}
// u32 assembled from 4 consecutive bytes at byte offset `b` (handles non-u32 alignment).
uint u32_at(uint base, uint b) {
    return byte_u(base, b) | (byte_u(base, b + 1u) << 8u)
         | (byte_u(base, b + 2u) << 16u) | (byte_u(base, b + 3u) << 24u);
}

void main() {
    uint n = gl_GlobalInvocationID.x;
    if (n >= nout) {
        return;
    }
    uint nblocks = k / 256u;
    uint rowbase = n * nblocks * BLK_U32;
    float acc = 0.0;
    for (uint blk = 0u; blk < nblocks; blk++) {
        uint base = rowbase + blk * BLK_U32;
        float d = float(unpackHalf2x16(w[base]).x);   // d in low half of u32[0] (byte 0..2)
        uint xblk = blk * 256u;
        for (uint e = 0u; e < 8u; e++) {
            uint aux0 = u32_at(base, 2u + 8u * e);
            uint aux1 = u32_at(base, 2u + 8u * e + 4u);
            float db = d * (0.5 + float(aux1 >> 28u)) * 0.25;
            for (uint l = 0u; l < 4u; l++) {
                uint gi = (aux0 >> (8u * l)) & 0xFFu;
                uint glo = IQ2XXS_GRID[2u * gi];
                uint ghi = IQ2XXS_GRID[2u * gi + 1u];
                uint signs = KSIGNS[(aux1 >> (7u * l)) & 127u];
                uint xoff = xblk + e * 32u + l * 8u;
                for (uint j = 0u; j < 8u; j++) {
                    uint g = (j < 4u) ? ((glo >> (8u * j)) & 0xFFu)
                                      : ((ghi >> (8u * (j - 4u))) & 0xFFu);
                    float sign = ((signs >> j) & 1u) != 0u ? -1.0 : 1.0;
                    acc += db * float(g) * sign * x[xoff + j];
                }
            }
        }
    }
    y[n] = acc;
}