ai.onnx.LayerNormalization / build /webgpu /layer-normalization.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
767d527 verified
Raw History Blame
11.7 kB
{% macro offset_fn(fn_name, opShape, opRank, op_same, op_numel, outShape, outRank, out_numel) %}
fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif %}) -> u32 {
{% if out_numel == 0 %}
return 0u;
{% elif op_numel == 1 %}
return 0u;
{% elif op_same %}
return out_index;
{% else %}
var offset = 0u;
{% for axis in range(outRank) %}
{% set op_axis = axis - (outRank - opRank) %}
{% if op_axis >= 0 and opShape[op_axis] != 1 %}
{% set c_stride = namespace(value=1) %}
{% for j in range(axis + 1, outRank) %}
{% set c_stride.value = c_stride.value * outShape[j] %}
{% endfor %}
{% set op_stride = namespace(value=1) %}
{% for j in range(op_axis + 1, opRank) %}
{% set op_stride.value = op_stride.value * opShape[j] %}
{% endfor %}
{% if c_stride.value == 1 %}
let coord{{ axis }} = out_index % {{ outShape[axis] }}u;
{% else %}
let coord{{ axis }} = (out_index / {{ c_stride.value }}u) % {{ outShape[axis] }}u;
{% endif %}
{% if op_stride.value == 1 %}
offset = offset + coord{{ axis }};
{% else %}
offset = offset + coord{{ axis }} * {{ op_stride.value }}u;
{% endif %}
{% endif %}
{% endfor %}
return offset;
{% endif %}
}{% endmacro %}
{% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
{% set op_numel = namespace(value=1) %}
{% for d in opShape %}
{% set op_numel.value = op_numel.value * d %}
{% endfor %}
{% set out_numel = namespace(value=1) %}
{% for d in outShape %}
{% set out_numel.value = out_numel.value * d %}
{% endfor %}
{{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %}){% endmacro %}
{{ env.wgsl.resourceDeclarations }}
const HIDDEN: u32 = {{ hiddenSize }}u;
const EPSILON: f32 = {{ epsilon }};
const WG: u32 = {{ workgroupSize }}u;
{% set xNumel = namespace(value=1) %}
{% for dim in xShape %}
{% set xNumel.value = xNumel.value * dim %}
{% endfor %}
{% set scaleNumel = namespace(value=1) %}
{% for dim in scaleShape %}
{% set scaleNumel.value = scaleNumel.value * dim %}
{% endfor %}
{% if scaleNumel.value != 1 %}
{{ offset_fn("scale_offset", scaleShape, scaleShape | length, scaleShape == xShape, scaleNumel.value, xShape, xShape | length, xNumel.value) }}
{% endif %}
{% if hasBias %}
{% set biasNumel = namespace(value=1) %}
{% for dim in biasShape %}
{% set biasNumel.value = biasNumel.value * dim %}
{% endfor %}
{% if biasNumel.value != 1 %}
{{ offset_fn("bias_offset", biasShape, biasShape | length, biasShape == xShape, biasNumel.value, xShape, xShape | length, xNumel.value) }}
{% endif %}
{% endif %}
{% if scalar == "f16" %}
{% set modeSpec = "layer" %}
{% set halfWriteMean = writeMean %}
{% set halfWriteInv = writeInvStdDev %}
{% set halfScaleOffset = "0u" if scaleNumel.value == 1 else broadcast_offset_call("scale_offset", scaleShape, xShape, "base + i") %}
{% if hasBias %}
{% set halfBiasOffset = "0u" if biasNumel.value == 1 else broadcast_offset_call("bias_offset", biasShape, xShape, "base + i") %}
{% endif %}
{% set halfOutputScalar = "f16" %}
fn round_f16_bits_rte(value: f32) -> u32 {
let bits = bitcast<u32>(value);
let sign = (bits >> 16u) & 0x8000u;
let exponent_f32 = (bits >> 23u) & 0xffu;
let mantissa_f32 = bits & 0x7fffffu;
if (exponent_f32 == 0xffu) {
if (mantissa_f32 != 0u) {
return 0x7e00u;
}
return sign | 0x7c00u;
}
var exponent_f16 = i32(exponent_f32) - 127 + 15;
if (exponent_f16 >= 0x1f) {
return sign | 0x7c00u;
}
if (exponent_f16 <= 0) {
if (exponent_f16 < -10) {
return sign;
}
let significand = mantissa_f32 | 0x800000u;
let shift = u32(14 - exponent_f16);
let halfway = 1u << (shift - 1u);
let discarded = significand & ((1u << shift) - 1u);
var fraction = significand >> shift;
if (discarded > halfway || (discarded == halfway && (fraction & 1u) == 1u)) {
fraction = fraction + 1u;
}
return sign | fraction;
}
let halfway = 1u << 12u;
let discarded = mantissa_f32 & 0x1fffu;
var mantissa_f16 = mantissa_f32 >> 13u;
if (discarded > halfway || (discarded == halfway && (mantissa_f16 & 1u) == 1u)) {
mantissa_f16 = mantissa_f16 + 1u;
if (mantissa_f16 == 0x400u) {
mantissa_f16 = 0u;
exponent_f16 = exponent_f16 + 1;
}
}
if (exponent_f16 >= 0x1f) {
return sign | 0x7c00u;
}
return sign | (u32(exponent_f16) << 10u) | mantissa_f16;
}
fn widen_f16_bits(value: u32) -> f32 {
return unpack2x16float(value & 0xffffu).x;
}
// Typed ONNX edges must survive arithmetic fusion and narrow/wide casts.
// Integer rounding also fixes the ties-to-even rule independently of the
// implementation's floating-point conversion rounding mode.
fn half_stage(value: f32) -> f32 {
return widen_f16_bits(round_f16_bits_rte(value));
}
// Half output magnifies statistics errors at rounding midpoints. Keep a low
// residual through the reduction and normalization, then round the typed
// float32-normalized/half-scale/half-bias edges explicitly. This is the standard ONNX half path;
// other normalization contracts retain their existing arithmetic.
{% set halfWriteMean = halfWriteMean if halfWriteMean is defined else (writeStats and modeSpec == "layer") %}
{% set halfWriteInv = halfWriteInv if halfWriteInv is defined else writeStats %}
fn pair_add(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> {
let s = fma(a.x, 1.0, b.x);
// Materialize each rounded subtraction in the error-free transform. Plain
// cancellation expressions do not preserve the intended evaluation tree on
// every shader backend.
let bv = fma(-1.0, a.x, s);
let av = fma(-1.0, bv, s);
let a_error = fma(-1.0, av, a.x);
let b_error = fma(-1.0, bv, b.x);
let error = fma(a_error, 1.0, b_error);
let e = fma(fma(error, 1.0, a.y), 1.0, b.y);
let hi = fma(s, 1.0, e);
return vec2<f32>(hi, fma(-1.0, fma(-1.0, s, hi), e));
}
fn pair_mul(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> {
let p = fma(a.x, b.x, 0.0);
let error = fma(a.x, b.x, -p);
return pair_add(vec2<f32>(p, 0.0), vec2<f32>(error + (a.x * b.y + b.x * a.y), 0.0));
}
fn pair_div(a: vec2<f32>, b: f32) -> vec2<f32> {
let q = a.x / b;
let residual = pair_add(a, -pair_mul(vec2<f32>(q, 0.0), vec2<f32>(b, 0.0)));
return pair_add(vec2<f32>(q, 0.0), vec2<f32>((residual.x + residual.y) / b, 0.0));
}
fn pair_inverse_sqrt(a: vec2<f32>) -> vec2<f32> {
let r = inverseSqrt(a.x);
let rr = pair_mul(vec2<f32>(r, 0.0), vec2<f32>(r, 0.0));
let residual = pair_add(vec2<f32>(1.0, 0.0), -pair_mul(a, rr));
return pair_add(vec2<f32>(r, 0.0), vec2<f32>((0.5 * r) * (residual.x + residual.y), 0.0));
}
fn half_normalized(value: vec2<f32>) -> f32 {
// stash_type=1 materializes Normalized as float32 before its cast to half.
// Collapse the compensated residual at that typed edge; rounding the pair
// directly to half can choose a different result at a float32 midpoint.
return widen_f16_bits(round_f16_bits_rte(fma(value.x, 1.0, value.y)));
}
var<workgroup> partial: array<vec2<f32>, WG>;
fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
partial[tid] = value;
workgroupBarrier();
for (var step = WG / 2u; step > 0u; step /= 2u) {
if (tid < step) { partial[tid] = pair_add(partial[tid], partial[tid + step]); }
workgroupBarrier();
}
let total = partial[0];
workgroupBarrier();
return total;
}
{% set reduceArgs = "tid" %}
fn load_value(index: u32) -> f32 {
return f32(x[index]);
}
fn normalize_half_row(row: u32, tid: u32
) {
if (row >= params.rows) { return; }
let base = row * HIDDEN;
var local_sum = vec2<f32>(0.0);
for (var i = tid; i < HIDDEN; i += WG) {
local_sum = pair_add(local_sum, vec2<f32>(load_value(base + i), 0.0));
}
let mean = pair_div(reduce_pair(local_sum, {{ reduceArgs }}), f32(HIDDEN));
var local_square = vec2<f32>(0.0);
for (var i = tid; i < HIDDEN; i += WG) {
let centered = pair_add(vec2<f32>(load_value(base + i), 0.0), -mean);
local_square = pair_add(local_square, pair_mul(centered, centered));
}
let variance = pair_div(reduce_pair(local_square, {{ reduceArgs }}), f32(HIDDEN));
let inv = pair_inverse_sqrt(pair_add(variance, vec2<f32>(EPSILON, 0.0)));
{% if halfWriteMean %}
if (tid == 0u) { mean_out[row] = mean.x; }
{% endif %}
{% if halfWriteInv %}
if (tid == 0u) { inv_std_out[row] = inv.x; }
{% endif %}
for (var i = tid; i < HIDDEN; i += WG) {
let normalized = half_normalized(pair_mul(pair_add(vec2<f32>(load_value(base + i), 0.0), -mean), inv));
var value = half_stage(normalized * f32(scale[{{ halfScaleOffset | default("i") }}]));
{% if modeSpec == "layer" and hasBias %}
value = half_stage(value + f32(bias[{{ halfBiasOffset | default("i") }}]));
{% endif %}
y[base + i] = {{ halfOutputScalar }}(value);
}
}
@compute @workgroup_size(WG, 1, 1)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
let row = wg.x + wg.y * params.rowStride;
normalize_half_row(row, lid.x);
}
{% else %}
var<workgroup> partial: array<f32, WG>;
var<workgroup> row_mean: f32;
var<workgroup> row_inv: f32;
{% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
{% if op == "max" or op == "min" %}
{{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
{{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ 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, reuse=false) %}
var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
loop {
if ({{ svar }} == 0u) {
break;
}
if ({{ idx }} < {{ svar }}) {
{% for a in arrays %}
{{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
{% endfor %}
}
{{ svar }} = {{ svar }} / 2u;
workgroupBarrier();
}{% endmacro %}
// Reusing partial after this reduction requires a barrier between the read of
// partial[0] and the next write, or the next round can race the prior readers.
fn reduce_sum(value: f32, tid: u32) -> f32 {
partial[tid] = value;
workgroupBarrier();
{{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
return partial[0];
}
@compute @workgroup_size(WG, 1, 1)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
let row = wg.x + wg.y * params.rowStride;
if (row >= params.rows) {
return;
}
let tid = lid.x;
let base = row * HIDDEN;
var local_sum = 0.0;
for (var d = tid; d < HIDDEN; d = d + WG) {
let value = f32(x[base + d]);
local_sum = local_sum + value;
}
let sum = reduce_sum(local_sum, tid);
if (tid == 0u) {
row_mean = sum / f32(HIDDEN);
}
workgroupBarrier();
var local_var_sum = 0.0;
for (var d = tid; d < HIDDEN; d = d + WG) {
let diff = f32(x[base + d]) - row_mean;
local_var_sum = local_var_sum + diff * diff;
}
let var_sum = reduce_sum(local_var_sum, tid);
if (tid == 0u) {
let variance = var_sum / f32(HIDDEN);
row_inv = inverseSqrt(variance + EPSILON);
{% if writeMean %}
mean_out[row] = row_mean;
{% endif %}
{% if writeInvStdDev %}
inv_std_out[row] = row_inv;
{% endif %}
}
workgroupBarrier();
for (var d = tid; d < HIDDEN; d = d + WG) {
let index = base + d;
let normalized = (f32(x[index]) - row_mean) * row_inv;
var value = normalized * f32(scale[{% if scaleNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("scale_offset", scaleShape, xShape, "index") }}{% endif %}]);
{% if hasBias %}
value = value + f32(bias[{% if biasNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("bias_offset", biasShape, xShape, "index") }}{% endif %}]);
{% endif %}
y[index] = {{ scalar }}(value);
}
}
{% endif %}