com.microsoft.SkipLayerNormalization / build /webgpu /norm-skip-row.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 2e7068faf55e
2760d09 verified
Raw
History Blame
7.35 kB
{% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
{% if op == "max" %}
{{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
{%- else %}
{{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
{%- endif %}
{% endmacro %}
{% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
loop {
{% if form == "head" %}
{% if breakInline %}
if ({{ svar }} == 0u) { break; }
{% else %}
if ({{ svar }} == 0u) {
break;
}
{% endif %}
{% endif %}
{% if bodyInline %}
if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
{% else %}
if ({{ idx }} < {{ svar }}) {
{% for a in arrays %}
{{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
{% endfor %}
}
{% endif %}
{% if form == "head" %}
{% if barrierFirst %}
workgroupBarrier();
{{ svar }} = {{ svar }} / 2u;
{% else %}
{{ svar }} = {{ svar }} / 2u;
workgroupBarrier();
{% endif %}
{% else %}
workgroupBarrier();
if ({{ svar }} == 1u) {
break;
}
{{ svar }} = {{ svar }} / 2u;
{% endif %}
}
{%- endmacro %}
/* One workgroup normalizes each row of residual = input + skip, with an
* optional bias. */
{% set degenerateRow = (not simplified) and hiddenSize == 1 %}
{% if useSubgroups and not degenerateRow %}
enable subgroups;
{% endif %}
{{ env.wgsl.resourceDeclarations }}
{% if not degenerateRow or writeResidualSum %}
const HIDDEN: u32 = {{ hiddenSize }}u;
{% endif %}
const WG: u32 = {{ workgroupSize }}u;
{% if simplified %}
var<workgroup> partial: array<f32, WG>;
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
fn {{ name }}(value: f32, tid: u32) -> f32 {
{{ 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) {
{% if mode == "max" %}
{{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
{% else %}
{{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
{% endif %}
}
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];
{% if trailingBarrier %}
workgroupBarrier();
{% endif %}
return reduced;
}
{% endmacro %}
{{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
var<workgroup> row_inv: f32;
{% else %}
{% if not degenerateRow %}
var<workgroup> pair_partial: array<vec2<f32>, WG>;
{% if useSubgroups %}
fn reduce_pair(value: vec2<f32>, sg_lane: u32, sg_id: u32, num_sg: u32) -> vec2<f32> {
let s = vec2<f32>(subgroupAdd(value.x), subgroupAdd(value.y));
if (num_sg == 1u) {
return s;
}
if (sg_lane == 0u) {
pair_partial[sg_id] = s;
}
workgroupBarrier();
var total = vec2<f32>(0.0, 0.0);
for (var i = 0u; i < num_sg; i = i + 1u) {
total = total + pair_partial[i];
}
return total;
}
{% else %}
fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
pair_partial[tid] = value;
workgroupBarrier();
{{ wgsl_tree_fold(["pair_partial"], idx="tid", wg="WG", form="head") }}
return pair_partial[0];
}
{% endif %}
{% endif %}
{% endif %}
{% if not degenerateRow or writeResidualSum %}
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;
}
{% endif %}
@compute @workgroup_size(WG, 1, 1)
fn main(
@builtin(workgroup_id) wg: vec3<u32>,
@builtin(num_workgroups) nwg: vec3<u32>{% if not degenerateRow %},
@builtin(local_invocation_id) lid: vec3<u32>{% endif %}{% if useSubgroups and not degenerateRow %},
@builtin(subgroup_invocation_id) sg_lane: u32,
@builtin(subgroup_id) sg_id: u32,
@builtin(num_subgroups) num_sg: u32{% endif %}
) {
// 2D-folded row index: wg.y carries the high bits past the maxComputeWorkgroupsPerDimension
// workgroup-per-dimension dispatch limit. Reduces to wg.x when nwg.y == 1;
// the row >= params.rows guard drops the over-dispatched tail.
let row = wg.x + wg.y * nwg.x;
if (row >= params.rows) {
return;
}
{% if not degenerateRow %}
let tid = lid.x;
{% endif %}
{% if simplified %}
// 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);
}
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]));
}
{% elif degenerateRow %}
// HIDDEN == 1: the row's mean is its only element, so the centered value and
// the variance are exactly zero and the output reduces to beta. The closed
// form avoids computing that zero by subtracting two equal rounded values.
let row_inv = inverseSqrt(params.epsilon);
{% if writeResidualSum %}
let residual = residual_value(row, 0u);
input_skip_bias_sum[row] = {{ scalar }}(residual);
{% endif %}
// 0.0 * row_inv keeps the IEEE result when epsilon == 0 makes row_inv +Inf.
output[row] = {{ scalar }}(0.0 * row_inv * f32(gamma[0]){% if hasBeta %} + f32(beta[0]){% endif %});
{% else %}
// Shifted moments: accumulating (x - x[0], (x - x[0])^2) keeps the sums
// small for rows with a large common offset; every thread reconstructs the
// row mean and variance from the merged pair.
let shift = residual_value(row, 0u);
var acc = vec2<f32>(0.0, 0.0);
for (var d = tid; d < HIDDEN; d = d + WG) {
let centered = residual_value(row, d) - shift;
acc.x = acc.x + centered;
acc.y = acc.y + centered * centered;
}
{% if useSubgroups %}
let totals = reduce_pair(acc, sg_lane, sg_id, num_sg);
{% else %}
let totals = reduce_pair(acc, tid);
{% endif %}
let mean_delta = totals.x / f32(HIDDEN);
let row_mean = shift + mean_delta;
let variance = max(totals.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0);
let row_inv = inverseSqrt(variance + params.epsilon);
for (var d = 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_mean) * row_inv * f32(gamma[d]){% if hasBeta %} + f32(beta[d]){% endif %});
}
{% endif %}
}