hanzo-ml 0.11.76

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
// IQ3_XXS matrix-vector (decode). GGUF IQ3_XXS super-block = 98 bytes: d (f16, byte 0..2) + qs[64]
// (grid indices, byte 2..66) + ss[32] (scales+signs, byte 66..98). 98 not u32-aligned -> repack to
// a padded 100-byte / 25-u32 stride. Per 32-weight sub-block e: aux32 = ss[4e..4e+4]; db =
// d*(0.5+(aux32>>28))*0.5; sign byte = KSIGNS[(aux32>>7l)&127] (group l=lane>>3). element m=lane&7:
// half=m>=4, jj=m&3; grid index = qs[8e+2l+half] -> 4-byte IQ3XXS_GRID u32 entry; g = byte jj;
// sign = bit m. weight=db*g*(+/-1). Byte-exact with BlockIQ3xxs::to_float. One invocation per 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[]; };  // IQ3_XXS blocks, 25 u32 / 256-block
layout(set = 0, binding = 1) readonly buffer X { float  x[]; };
layout(set = 0, binding = 2) writeonly buffer Y { float y[]; };
layout(push_constant) uniform Pc { uint nout; uint k; };

const uint BLK_U32 = 25u;   // 100 bytes / 4 (98 real + 2 pad)

// IQ3XXS_GRID: 256 u32 entries (4 int8 each); KSIGNS: 128 bytes
const uint IQ3XXS_GRID[256] = uint[256](
    0x04040404u, 0x04040414u, 0x04040424u, 0x04040c0cu, 0x04040c1cu, 0x04040c3eu, 0x04041404u, 0x04041414u, 0x04041c0cu,
    0x04042414u, 0x04043e1cu, 0x04043e2cu, 0x040c040cu, 0x040c041cu, 0x040c0c04u, 0x040c0c14u, 0x040c140cu, 0x040c142cu,
    0x040c1c04u, 0x040c1c14u, 0x040c240cu, 0x040c2c24u, 0x040c3e04u, 0x04140404u, 0x04140414u, 0x04140424u, 0x04140c0cu,
    0x04141404u, 0x04141414u, 0x04141c0cu, 0x04141c1cu, 0x04141c3eu, 0x04142c0cu, 0x04142c3eu, 0x04143e2cu, 0x041c040cu,
    0x041c043eu, 0x041c0c04u, 0x041c0c14u, 0x041c142cu, 0x041c3e04u, 0x04240c1cu, 0x04241c3eu, 0x04242424u, 0x04242c3eu,
    0x04243e1cu, 0x04243e2cu, 0x042c040cu, 0x042c043eu, 0x042c1c14u, 0x042c2c14u, 0x04341c2cu, 0x04343424u, 0x043e0c04u,
    0x043e0c24u, 0x043e0c34u, 0x043e241cu, 0x043e340cu, 0x0c04040cu, 0x0c04041cu, 0x0c040c04u, 0x0c040c14u, 0x0c04140cu,
    0x0c04141cu, 0x0c041c04u, 0x0c041c14u, 0x0c041c24u, 0x0c04243eu, 0x0c042c04u, 0x0c0c0404u, 0x0c0c0414u, 0x0c0c0c0cu,
    0x0c0c1404u, 0x0c0c1414u, 0x0c14040cu, 0x0c14041cu, 0x0c140c04u, 0x0c140c14u, 0x0c14140cu, 0x0c141c04u, 0x0c143e14u,
    0x0c1c0404u, 0x0c1c0414u, 0x0c1c1404u, 0x0c1c1c0cu, 0x0c1c2434u, 0x0c1c3434u, 0x0c24040cu, 0x0c24042cu, 0x0c242c04u,
    0x0c2c1404u, 0x0c2c1424u, 0x0c2c2434u, 0x0c2c3e0cu, 0x0c34042cu, 0x0c3e1414u, 0x0c3e2404u, 0x14040404u, 0x14040414u,
    0x14040c0cu, 0x14040c1cu, 0x14041404u, 0x14041414u, 0x14041434u, 0x14041c0cu, 0x14042414u, 0x140c040cu, 0x140c041cu,
    0x140c042cu, 0x140c0c04u, 0x140c0c14u, 0x140c140cu, 0x140c1c04u, 0x140c341cu, 0x140c343eu, 0x140c3e04u, 0x14140404u,
    0x14140414u, 0x14140c0cu, 0x14140c3eu, 0x14141404u, 0x14141414u, 0x14141c3eu, 0x14142404u, 0x14142c2cu, 0x141c040cu,
    0x141c0c04u, 0x141c0c24u, 0x141c3e04u, 0x141c3e24u, 0x14241c2cu, 0x14242c1cu, 0x142c041cu, 0x142c143eu, 0x142c240cu,
    0x142c3e24u, 0x143e040cu, 0x143e041cu, 0x143e0c34u, 0x143e242cu, 0x1c04040cu, 0x1c040c04u, 0x1c040c14u, 0x1c04140cu,
    0x1c04141cu, 0x1c042c04u, 0x1c04342cu, 0x1c043e14u, 0x1c0c0404u, 0x1c0c0414u, 0x1c0c1404u, 0x1c0c1c0cu, 0x1c0c2424u,
    0x1c0c2434u, 0x1c14040cu, 0x1c14041cu, 0x1c140c04u, 0x1c14142cu, 0x1c142c14u, 0x1c143e14u, 0x1c1c0c0cu, 0x1c1c1c1cu,
    0x1c241c04u, 0x1c24243eu, 0x1c243e14u, 0x1c2c0404u, 0x1c2c0434u, 0x1c2c1414u, 0x1c2c2c2cu, 0x1c340c24u, 0x1c341c34u,
    0x1c34341cu, 0x1c3e1c1cu, 0x1c3e3404u, 0x24040424u, 0x24040c3eu, 0x24041c2cu, 0x24041c3eu, 0x24042c1cu, 0x24042c3eu,
    0x240c3e24u, 0x24141404u, 0x24141c3eu, 0x24142404u, 0x24143404u, 0x24143434u, 0x241c043eu, 0x241c242cu, 0x24240424u,
    0x24242c0cu, 0x24243424u, 0x242c142cu, 0x242c241cu, 0x242c3e04u, 0x243e042cu, 0x243e0c04u, 0x243e0c14u, 0x243e1c04u,
    0x2c040c14u, 0x2c04240cu, 0x2c043e04u, 0x2c0c0404u, 0x2c0c0434u, 0x2c0c1434u, 0x2c0c2c2cu, 0x2c140c24u, 0x2c141c14u,
    0x2c143e14u, 0x2c1c0414u, 0x2c1c2c1cu, 0x2c240c04u, 0x2c24141cu, 0x2c24143eu, 0x2c243e14u, 0x2c2c0414u, 0x2c2c1c0cu,
    0x2c342c04u, 0x2c3e1424u, 0x2c3e2414u, 0x34041424u, 0x34042424u, 0x34042434u, 0x34043424u, 0x340c140cu, 0x340c340cu,
    0x34140c3eu, 0x34143424u, 0x341c1c04u, 0x341c1c34u, 0x34242424u, 0x342c042cu, 0x342c2c14u, 0x34341c1cu, 0x343e041cu,
    0x343e140cu, 0x3e04041cu, 0x3e04042cu, 0x3e04043eu, 0x3e040c04u, 0x3e041c14u, 0x3e042c14u, 0x3e0c1434u, 0x3e0c2404u,
    0x3e140c14u, 0x3e14242cu, 0x3e142c14u, 0x3e1c0404u, 0x3e1c0c2cu, 0x3e1c1c1cu, 0x3e1c3404u, 0x3e24140cu, 0x3e24240cu,
    0x3e2c0404u, 0x3e2c0414u, 0x3e2c1424u, 0x3e341c04u
);

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
);


uint byte_u(uint base, uint b) {
    return (w[base + (b >> 2u)] >> ((b & 3u) * 8u)) & 0xFFu;
}
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);
        uint xblk = blk * 256u;
        for (uint e = 0u; e < 8u; e++) {
            uint aux32 = u32_at(base, 66u + 4u * e);   // ss[4e..4e+4]
            float db = d * (0.5 + float(aux32 >> 28u)) * 0.5;
            for (uint l = 0u; l < 4u; l++) {
                uint signs = KSIGNS[(aux32 >> (7u * l)) & 127u];
                uint xoff = xblk + e * 32u + l * 8u;
                for (uint m = 0u; m < 8u; m++) {
                    uint hf = (m < 4u) ? 0u : 1u;
                    uint jj = m & 3u;
                    uint gidx = byte_u(base, 2u + 8u * e + 2u * l + hf);  // qs[8e+2l+half]
                    uint entry = IQ3XXS_GRID[gidx];
                    uint g = (entry >> (8u * jj)) & 0xFFu;
                    float sign = ((signs >> m) & 1u) != 0u ? -1.0 : 1.0;
                    acc += db * float(g) * sign * x[xoff + m];
                }
            }
        }
    }
    y[n] = acc;
}