// COMPILE-VERIFIED on macOS (glslc 2026.2, target Vulkan 1.1 / SPIR-V 1.3). // Wave-4-B parity port: the in-place Hadamard butterfly is now executed by // 32 threads cooperatively (mirrors metal/polar.metal). The (1/QK_POLAR) // compensation is folded into the final per-row scalar so the parallel // multiply pass over buf[] is gone. // // PolarQuant 4-bit block (block_q4_polar, 82 bytes): // fp16 d // [0..1] per-block L2 norm // uchar qs[64] // [2..65] 4-bit codes, low nibble first // uchar qjl[16] // [66..81] optional 1-bit QJL residual // // Bits/element @ head_dim=128: with QJL = 5.125 bpw, without = 4.125 bpw. // // Ports `kernel_mul_mv_q4_polar_f32` from metal/polar.metal. // Decode steps (mirror dequantize_row_q4_polar_ref): // 1. Unpack 4-bit codes -> centroid LUT lookup (16 entries, Lloyd-Max N(0,1)). // 2. Optional QJL residual: 1 sign-bit applied to a deterministic +/-1 sign // vector (xorshift32 seeded with POLAR_QJL_SEED=42), magnitude // 0.5 / sqrt(QK_POLAR). // 3. In-place 128-element Walsh-Hadamard butterfly (7 stages). // 4. Compensate by 1/QK_POLAR (orthonormal-inverse Hadamard). // 5. Per-block L2 rescale by stored fp16 norm. #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; // k_blocks: (n_rows) packed 82-byte block_q4_polar, indexed via raw uint stream. layout(std430, binding = 0) readonly buffer KBlocks { uint k_packed[]; }; layout(std430, binding = 1) readonly buffer Q { float q[]; // (head_dim) }; layout(std430, binding = 2) writeonly buffer YOut { float y[]; // (n_rows) }; 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 the 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 QJL_RESIDUAL_BYTES = 16u; // QK_POLAR / 8 const uint POLAR_BLOCK_BYTES = 82u; // 2 + 64 + 16 const float POLAR_QJL_CORRECTION_MAGNITUDE = 0.5; // 1 / sqrt(128). const float POLAR_QJL_INV_SQRT_QK = 0.08838834764831845; 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 a tid==0 recurrent fill in the residual 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 ); 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 turbo3.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); } shared float buf[128]; // 1 block of decoded floats 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]; } // Threadgroup-cooperative 128-element Walsh-Hadamard butterfly. Mirrors the // Wave-4-B Metal port (polar_hadamard_inplace_tg32 in metal/polar.metal): // 32 threads × 2 of 64 (a+b, a-b) butterfly pairs per stage, one barrier() // between stages. Within a single stage every index 0..127 is touched by // exactly one pair, so reads and writes do not race. Caller MUST barrier() // before invoking so the input fill is visible to all threads. void polar_hadamard_inplace_tg32(uint tid) { for (uint h = 1u; h < QK_POLAR; h <<= 1) { uint p0 = tid; // 0..31 uint p1 = tid + 32u; // 32..63 uint twoh = h << 1; uint b0 = (p0 / h) * twoh; uint o0 = p0 - (p0 / h) * h; // p0 % h, branchless uint b1 = (p1 / h) * twoh; uint o1 = p1 - (p1 / h) * h; uint j0 = b0 + o0; uint j1 = b1 + o1; float a0 = buf[j0]; float c0 = buf[j0 + h]; float a1 = buf[j1]; float c1 = buf[j1 + h]; buf[j0] = a0 + c0; buf[j0 + h] = a0 - c0; buf[j1] = a1 + c1; buf[j1 + h] = a1 - c1; barrier(); } } 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; // Step 1+2: unpack 4-bit codes to centroids. 32 threads x 2 bytes each // covers all 64 bytes / 128 elements. for (uint b = tid; b < QK_POLAR / 2u; b += 32u) { uint byte = read_byte(blk_off + 2u + b); buf[2u * b] = POLAR_Q4_CENTROIDS[byte & 0x0Fu]; buf[2u * b + 1u] = POLAR_Q4_CENTROIDS[(byte >> 4) & 0x0Fu]; } barrier(); // Step 3: optional QJL residual. The xorshift32(seed=42) sign vector is a // literal constant table, so all 32 threads can apply it directly. if (push.use_qjl != 0u) { uint bit = read_byte(blk_off + 2u + 64u) & 1u; float sign_v = bit != 0u ? 1.0 : -1.0; float mag = POLAR_QJL_CORRECTION_MAGNITUDE * POLAR_QJL_INV_SQRT_QK; float scaled = sign_v * mag; for (uint i = tid; i < QK_POLAR; i += 32u) { buf[i] += scaled * POLAR_QJL_SIGNS[i]; } barrier(); } // Step 4: inverse Hadamard — threadgroup-cooperative 32-thread butterfly. // Replaces the previous tid==0 sequential 7-stage loop that was the // dominant cost in the Metal polar kernel (12.5× speedup on M4 Max). polar_hadamard_inplace_tg32(tid); // Step 5: dot product against q[]. Fold the (1/QK_POLAR) Hadamard // compensation and per-block L2 norm into one final scalar applied // after the tree reduction — saves a parallel multiply pass over buf[]. float acc = 0.0; for (uint i = tid; i < QK_POLAR; i += 32u) { acc += buf[i] * q[push.q_offset + i]; } 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; } }