// DRAFT: COMPILE-VERIFIED on Linux. Sister kernel to qjl.comp. // Ports `kernel_get_rows_qjl1_256` from metal/qjl.metal — decode one // QJL block to a fp32 head_dim vector using the projection matrix Pi. // // recon[i] = (||k|| * sqrt(pi/2) / proj_dim) * sum_j sign(j) * prj[i*proj_dim + j] // // Threadgroup size = 32. Each thread handles head_dim / 32 = 4 output rows. // Inner loop sums proj_dim=256 signed projection rows. #version 450 #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; layout(std430, binding = 0) readonly buffer KPacked { uint k_packed[]; // single block: 32 bytes signs + 2 bytes norm }; layout(std430, binding = 1) readonly buffer Pi { float prj[]; // (head_dim, proj_dim) row-major }; layout(std430, binding = 2) writeonly buffer Out { float out_buf[]; // head_dim }; layout(push_constant) uniform Push { uint head_dim; // must equal 128 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 float QJL_SCORE_SCALE = 1.2533141373155003 / 256.0; uint read_byte(uint b) { uint w = k_packed[b >> 2]; return (w >> ((b & 3u) * 8u)) & 0xFFu; } uint read_u16(uint b) { return read_byte(b) | (read_byte(b + 1u) << 8u); } float qjl_bf16_to_fp32(uint b16) { return uintBitsToFloat(b16 << 16); } void main() { uint tid = gl_LocalInvocationID.x; if (push.head_dim != QJL_HEAD_DIM || push.proj_dim != QJL_PROJECTION_DIM) return; uint norm16 = read_u16(QJL_PACKED_BYTES); float scale = QJL_SCORE_SCALE * qjl_bf16_to_fp32(norm16); // Each thread walks output rows i with stride = 32; sums proj_dim signed projections. for (uint i = tid; i < QJL_HEAD_DIM; i += 32u) { float acc = 0.0; uint row_off = i * QJL_PROJECTION_DIM; for (uint j = 0u; j < QJL_PROJECTION_DIM; ++j) { uint bit = (read_byte(j >> 3) >> (j & 7u)) & 1u; acc += bit != 0u ? prj[row_off + j] : -prj[row_off + j]; } out_buf[i] = scale * acc; } }