// COMPILE/HARNESS VERIFIED: pre-Hadamard-query hot path for PolarQuant. // // This is the Vulkan sibling of metal/polar.metal's // kernel_mul_mv_q4_polar_preht_f32. It uses: // // dot(H*x, q) == dot(x, H*q) // // where H is the unnormalised 128-point Walsh-Hadamard transform used by the // decoder. The caller supplies q_preht = H*q. That removes the per-row shared // scratch and 7-stage Hadamard butterfly from the attention-score hot path. // // DO NOT route this kernel behind an existing raw-q Polar graph dispatch. The // q buffer must already contain H*q or the result is mathematically wrong. #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; layout(std430, binding = 0) readonly buffer KBlocks { uint k_packed[]; }; layout(std430, binding = 1) readonly buffer Q { float q_preht[]; }; layout(std430, binding = 2) writeonly buffer YOut { float y[]; }; layout(push_constant) uniform Push { uint n_rows; uint head_dim; // must equal QK_POLAR (128) uint use_qjl; // 0 / 1 uint k_offset_bytes; // graph dispatch base offset for selected KV head uint q_offset; // graph dispatch base offset in float elements uint y_offset; // graph dispatch base offset in float elements } push; const uint QK_POLAR = 128u; const uint POLAR_BLOCK_BYTES = 82u; // 2 + 64 + 16 const float POLAR_QJL_CORRECTION_MAGNITUDE = 0.5; const float POLAR_QJL_INV_SQRT_QK = 0.08838834764831845; // 1 / sqrt(128) const float POLAR_INV_QK = 1.0 / 128.0; // Bit-identical to POLAR_Q4_CENTROIDS in // packages/native-plugins/polarquant-cpu/include/polarquant/polar_centroids.h. const float POLAR_Q4_CENTROIDS[16] = float[16]( -2.754354807, -2.093562707, -1.643041510, -1.279739752, -0.962640978, -0.672392117, -0.397897103, -0.131757782, 0.131757782, 0.397897103, 0.672392117, 0.962640978, 1.279739752, 1.643041510, 2.093562707, 2.754354807 ); // xorshift32(seed=42) sign vector used by the optional Polar QJL residual. // Literal table avoids the recurrent xorshift chain in the hot path. const float POLAR_QJL_SIGNS[128] = float[128]( -1.0, -1.0, 1.0, -1.0, -1.0, -1.0, -1.0, -1.0, 1.0, 1.0, -1.0, 1.0, -1.0, -1.0, 1.0, 1.0, -1.0, -1.0, 1.0, -1.0, 1.0, 1.0, -1.0, -1.0, -1.0, 1.0, -1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, -1.0, 1.0, -1.0, 1.0, -1.0, 1.0, -1.0, -1.0, -1.0, -1.0, 1.0, -1.0, 1.0, 1.0, 1.0, 1.0, 1.0, -1.0, 1.0, -1.0, -1.0, 1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, 1.0, -1.0, 1.0, 1.0, -1.0, 1.0, -1.0, 1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, -1.0, 1.0, -1.0, -1.0, -1.0, 1.0, -1.0, -1.0, 1.0, 1.0, 1.0, 1.0, -1.0, -1.0, 1.0, -1.0, 1.0, 1.0, -1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, 1.0, 1.0, -1.0, 1.0, -1.0, 1.0, 1.0, -1.0, -1.0, 1.0, -1.0, -1.0, 1.0, -1.0, 1.0, -1.0, -1.0, 1.0, -1.0, -1.0, 1.0, 1.0, 1.0, -1.0, 1.0, -1.0, -1.0, 1.0 ); shared float reduce_scratch[32]; 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); } // fp16 -> fp32 (manual; same routine as polar.comp). float fp16_to_fp32(uint h16) { uint sign = (h16 & 0x8000u) << 16; uint exp = (h16 >> 10) & 0x1Fu; uint mant = h16 & 0x3FFu; uint u; if (exp == 0u) { if (mant == 0u) { u = sign; } else { uint e = 1u; while ((mant & 0x400u) == 0u) { mant <<= 1; e += 1u; } mant &= 0x3FFu; u = sign | ((127u - 15u - e + 1u) << 23) | (mant << 13); } } else if (exp == 0x1Fu) { u = sign | 0x7F800000u | (mant << 13); } else { u = sign | ((exp + 127u - 15u) << 23) | (mant << 13); } return uintBitsToFloat(u); } 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 row = gl_WorkGroupID.x; if (row >= push.n_rows || push.head_dim != QK_POLAR) return; uint blk_off = push.k_offset_bytes + row * POLAR_BLOCK_BYTES; float acc = 0.0; for (uint b = tid; b < QK_POLAR / 2u; b += 32u) { uint byte = read_byte(blk_off + 2u + b); uint i0 = 2u * b; uint i1 = i0 + 1u; float x0 = POLAR_Q4_CENTROIDS[byte & 0x0Fu]; float x1 = POLAR_Q4_CENTROIDS[(byte >> 4) & 0x0Fu]; if (push.use_qjl != 0u) { float mag = POLAR_QJL_CORRECTION_MAGNITUDE * POLAR_QJL_INV_SQRT_QK; uint bit = read_byte(blk_off + 2u + 64u) & 1u; float sign_v = bit != 0u ? 1.0 : -1.0; float scaled = sign_v * mag; x0 += scaled * POLAR_QJL_SIGNS[i0]; x1 += scaled * POLAR_QJL_SIGNS[i1]; } acc += x0 * q_preht[push.q_offset + i0]; acc += x1 * q_preht[push.q_offset + i1]; } float sum = reduce_sum_32(acc, tid); if (tid == 0u) { uint norm16 = read_u16(blk_off); float l2 = fp16_to_fp32(norm16); y[push.y_offset + row] = sum * l2 * POLAR_INV_QK; } }