// DRAFT: COMPILE-VERIFIED on Linux (glslc 2022.3, target Vulkan 1.1 / SPIR-V 1.3). // Hardware verification still required — see packages/inference/README.md // "Verification matrix". // // QJL = K-side compression: store sign(Pi*k) packed 8-per-byte plus per-token // bf16 norm. Q is sketched once via the same Pi and the score per // (q_head, token) is reconstructed from the packed signs. // // Block layout (block_qjl1_256, 34 bytes, alignment 2): // uchar qs[32] // 256 sign bits, LSB = bit 0 of byte 0 // ushort norm_bf16 // bf16 storage of ||k||_2 // // Total compressed bits/element at head_dim=128: 34*8 / 128 = 2.125 bpw, // 7.53x vs bf16 K-cache. // // Ports the three Metal kernels in metal/qjl.metal: // - kernel_attn_score_qjl1_256 : the attention-score hot path. // - kernel_get_rows_qjl1_256 : decode signs * Pi reconstruction (debug path). // - kernel_mul_mv_qjl1_256_f32 : matrix-vector against an fp32 sketch. // // SPIR-V cannot host three entrypoints in a single .comp easily and the // Vulkan wrapper that consumes these kernels expects one .spv per kernel, // so this file builds the SCORE entrypoint. Sister files // `qjl_get_rows.comp` and `qjl_mul_mv.comp` carry the other two. #version 450 #extension GL_EXT_shader_16bit_storage : require #extension GL_EXT_shader_explicit_arithmetic_types_int8 : require #extension GL_EXT_shader_explicit_arithmetic_types_int16 : require layout(local_size_x = 32, local_size_y = 1, local_size_z = 1) in; // q_sketch: (n_heads, proj_dim=256) fp32, host pre-projected. layout(std430, binding = 0) readonly buffer QSketch { float q_sketch[]; }; // packed_k: (n_kv_heads, n_tokens) tightly packed 34-byte block_qjl1_256. // Indexed as raw uint stream; std430 forces 4-byte alignment so the // 34-byte block is read by walking byte-by-byte. layout(std430, binding = 1) readonly buffer KPacked { uint k_packed[]; }; // Output: (n_heads, n_tokens) fp32 scores. layout(std430, binding = 2) writeonly buffer ScoresOut { float scores[]; }; layout(push_constant) uniform Push { uint n_heads; uint n_kv_heads; // n_heads / n_kv_heads = GQA factor (>= 1) uint n_tokens; uint proj_dim; // must equal 256 } push; const uint QJL_HEAD_DIM = 128u; const uint QJL_PROJECTION_DIM = 256u; const uint QJL_PACKED_BYTES = 32u; const uint QJL_BLOCK_BYTES = 34u; // sqrt(pi/2) / proj_dim — matches qjl_score_qk_ref's scl_base. const float QJL_SCORE_SCALE = 1.2533141373155003 / 256.0; // Read one byte from k_packed at absolute byte offset `b`. uint read_byte(uint b) { uint w = k_packed[b >> 2]; return (w >> ((b & 3u) * 8u)) & 0xFFu; } // Read a little-endian u16 at byte offset `b`. uint read_u16(uint b) { return read_byte(b) | (read_byte(b + 1u) << 8u); } // bf16 -> fp32 zero-extension (matches qjl_bf16_to_fp32 in qjl-cpu). float qjl_bf16_to_fp32(uint b16) { return uintBitsToFloat(b16 << 16); } // 32-thread tree reduction over a threadgroup-shared scratch. Driver-portable: // does NOT depend on any specific subgroup size, unlike subgroupAdd. shared float reduce_scratch[32]; float reduce_sum_32(float v, uint tid) { reduce_scratch[tid] = v; barrier(); for (uint stride = 16u; stride > 0u; stride >>= 1) { if (tid < stride) { reduce_scratch[tid] += reduce_scratch[tid + stride]; } barrier(); } return reduce_scratch[0]; } void main() { uint tid = gl_LocalInvocationID.x; uint h_q = gl_WorkGroupID.x; uint t = gl_WorkGroupID.y; if (h_q >= push.n_heads || t >= push.n_tokens) return; uint gqa = push.n_heads / push.n_kv_heads; uint h_k = h_q / gqa; // Byte offset of block (h_k, t) inside packed_k. uint blk_off = (h_k * push.n_tokens + t) * QJL_BLOCK_BYTES; // Each of 32 threads owns one byte (8 sign bits) of qs[]. Mirrors the // Wave-4-B Metal port (kernel_attn_score_qjl1_256): pack 8 contiguous // q_sketch floats into two vec4s, build branchless ±1 sign vectors via // ((bit << 1) - 1), and chain two fma()s. glslc on Vulkan 1.1 emits // OpVectorTimesScalar/OpFMul/OpExtInst Fma against vec4 — same pattern // as the Metal float4 fma the M4 Max kernel uses. uint byte_i = tid; uint bits = read_byte(blk_off + byte_i); uint base = byte_i * 8u; uint q_off = h_q * QJL_PROJECTION_DIM + base; vec4 q0 = vec4(q_sketch[q_off + 0u], q_sketch[q_off + 1u], q_sketch[q_off + 2u], q_sketch[q_off + 3u]); vec4 q1 = vec4(q_sketch[q_off + 4u], q_sketch[q_off + 5u], q_sketch[q_off + 6u], q_sketch[q_off + 7u]); vec4 s0 = vec4( float(int(((bits >> 0) & 1u) << 1) - 1), float(int(((bits >> 1) & 1u) << 1) - 1), float(int(((bits >> 2) & 1u) << 1) - 1), float(int(((bits >> 3) & 1u) << 1) - 1)); vec4 s1 = vec4( float(int(((bits >> 4) & 1u) << 1) - 1), float(int(((bits >> 5) & 1u) << 1) - 1), float(int(((bits >> 6) & 1u) << 1) - 1), float(int(((bits >> 7) & 1u) << 1) - 1)); vec4 acc4 = fma(q0, s0, fma(q1, s1, vec4(0.0))); float acc = acc4.x + acc4.y + acc4.z + acc4.w; float sum = reduce_sum_32(acc, tid); if (tid == 0u) { uint norm16 = read_u16(blk_off + QJL_PACKED_BYTES); float norm_k = qjl_bf16_to_fp32(norm16); scores[h_q * push.n_tokens + t] = QJL_SCORE_SCALE * norm_k * sum; } }