{% if usesF16 is defined and usesF16 %}enable f16; {% endif %}{{ env.wgsl.resourceDeclarations }} // One invocation per (batch, query head, query token) applies optional per-head // RMS normalization followed by NeoX half-split rotary embedding at the query's // absolute position. Present-cache keys are already transformed, so this pass // writes only Q in BSH layout. f16 queries are widened for normalization, // trigonometry, and rotation. The result narrows only when the following // attention pass consumes f16 operand tiles; other configurations store f32. const HEAD_DIM: u32 = {{ headDim }}u; const HEAD_DIM_V4: u32 = {{ headDim }}u / 4u; const Q_HEADS: u32 = {{ qHeads }}u; const Q_HIDDEN: u32 = {{ qHidden }}u; const Q_HIDDEN_V4: u32 = {{ qHidden }}u / 4u; const WG: u32 = {{ copyWorkgroupSize }}u; {% set HAS_QNORM = hasQNorm is not defined or hasQNorm %} const HALF: u32 = {{ half }}u; {% if HAS_QNORM %} const QK_EPS: f32 = {{ qkEps }}; {% endif %} @compute @workgroup_size(WG) fn main(@builtin(global_invocation_id) gid: vec3) { // 2D-folded flat index: gid.y carries the high bits past the per-axis dispatch fold width. // Reduces to gid.x when the dispatch does not fold. let qi = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG; let total = params.batch * Q_HEADS * params.qSeq; if (qi >= total) { return; } let s = qi % params.qSeq; let tmp = qi / params.qSeq; let h = tmp % Q_HEADS; let b = tmp / Q_HEADS; let absPos = params.totalSeq - params.qSeq + s; let base = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM; var q: array; for (var d = 0u; d < HEAD_DIM; d = d + 1u) { q[d] = f32(query[base + d]); } {% if HAS_QNORM %} var ss = 0.0; for (var d = 0u; d < HEAD_DIM; d = d + 1u) { ss = ss + q[d] * q[d]; } let invRms = inverseSqrt(ss / f32(HEAD_DIM) + QK_EPS); for (var d = 0u; d < HEAD_DIM; d = d + 1u) { q[d] = q[d] * invRms * f32(q_norm_weight[d]); } {% endif %} for (var d = 0u; d < HALF; d = d + 1u) { let cs = f32(cos_cache[absPos * HALF + d]); let sn = f32(sin_cache[absPos * HALF + d]); let x0 = q[d]; let x1 = q[d + HALF]; q[d] = x0 * cs - x1 * sn; q[d + HALF] = x1 * cs + x0 * sn; } let base4 = (b * params.qSeq + s) * Q_HIDDEN_V4 + h * HEAD_DIM_V4; for (var c = 0u; c < HEAD_DIM_V4; c = c + 1u) { qout[base4 + c] = {{ "vec4(" if (qPrepF16 is defined and qPrepF16) else "" }}vec4(q[c * 4u], q[c * 4u + 1u], q[c * 4u + 2u], q[c * 4u + 3u]){{ ")" if (qPrepF16 is defined and qPrepF16) else "" }}; } }