/* Normalize residual = input + skip, with an optional bias. Reductions use * one workgroup per row; closed-form one-element rows use one invocation. */ {% set degenerateRow = false %} {{ env.wgsl.resourceDeclarations }} const HIDDEN: u32 = {{ hiddenSize }}u; const WG: u32 = {{ workgroupSize }}u; var partial: array; {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %} fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} { {{ buffer }}[tid] = value; workgroupBarrier(); // Ceil-halving keeps every lane when the workgroup size is not a power of // two. For even n this matches the power-of-two tree order; for odd n, lanes // [0, n-half) fold the upper tail while the middle lane carries forward. var n: u32 = {{ wg }}; loop { let half = (n + 1u) / 2u; if (tid < n - half) { {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half]; } workgroupBarrier(); n = half; if (n == 1u) { break; } } // The default trailing barrier makes this helper safe for back-to-back calls: every lane reads // slot 0 here, so the next call's first store must not run until all lanes have read it. // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit. let reduced = {{ buffer }}[0]; workgroupBarrier(); return reduced; } {% endmacro %} {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }} var row_inv: f32; fn residual_value(row: u32, d: u32) -> f32 { let index = row * HIDDEN + d; var value = f32(input[index]) + f32(skip[index]); {% if hasBias %} value = value + f32(bias[d]); {% endif %} return value; } @compute @workgroup_size(WG, 1, 1) fn main( @builtin({{ "global_invocation_id" if degenerateRow else "workgroup_id" }}) {{ "gid" if degenerateRow else "wg" }}: vec3, @builtin(local_invocation_id) lid: vec3) { // Fold the row grid across workgroups; independent rows also include the // invocation offset. The bounds guard drops the final dispatch tail. let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u; if (row >= params.rows) { return; } let tid = lid.x; // RMS normalization uses one sum-of-squares sweep, without a mean or beta. var local_sq = 0.0; for (var d: u32 = tid; d < HIDDEN; d = d + WG) { let value = residual_value(row, d); local_sq = local_sq + value * value; } let sq = reduce_sum(local_sq, tid); if (tid == 0u) { row_inv = inverseSqrt(sq / f32(HIDDEN) + params.epsilon); {% if writeMean is defined and writeMean %} mean[row] = 0.0; {% endif %} {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %} inv_std_var[row] = row_inv; {% endif %} } workgroupBarrier(); for (var d: u32 = tid; d < HIDDEN; d = d + WG) { let index = row * HIDDEN + d; let residual = residual_value(row, d); {% if writeResidualSum %} input_skip_bias_sum[index] = {{ scalar }}(residual); {% endif %} output[index] = {{ scalar }}(residual * row_inv * f32(gamma[d])); } }