Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
5ec57d5 verified
Raw History Blame
3.14 kB
/* 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<workgroup> partial: array<f32, WG>;
{% 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<workgroup> 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<u32>,
@builtin(local_invocation_id) lid: vec3<u32>) {
// 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]));
}
}