com.microsoft.SkipLayerNormalization / build /webgpu /norm-skip-row-vec4.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 2e7068faf55e
2760d09 verified
Raw
History Blame
4.74 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 %}{% set useSubgroups = source.useSubgroups %}
{% if source.usesF16 %}
enable f16;
{% endif %}
{% if useSubgroups %}
enable subgroups;
{% endif %}
{{ env.wgsl.resourceDeclarations }}
const HIDDEN: u32 = {{ source.hidden }}u;
const HIDDEN_V: u32 = {{ source.hiddenVec }}u;
const WG: u32 = {{ source.wg }}u;
var<workgroup> sg_partials: array<vec2<f32>, WG>;
fn reduce_pair(value: vec2<f32>{% if useSubgroups %}, sg_lane: u32, sg_id: u32, num_sg: u32{% else %}, tid: u32{% endif %}) -> vec2<f32> {
{% if useSubgroups %}
let s = vec2<f32>(subgroupAdd(value.x), subgroupAdd(value.y));
if (num_sg == 1u) {
return s;
}
if (sg_lane == 0u) {
sg_partials[sg_id] = s;
}
workgroupBarrier();
var total = vec2<f32>(0.0, 0.0);
for (var i = 0u; i < num_sg; i = i + 1u) {
total = total + sg_partials[i];
}
return total;
{% else %}
// No-subgroup tier: workgroup barrier tree-reduction (WG is a power of two).
sg_partials[tid] = value;
workgroupBarrier();
{{ wgsl_tree_fold(["sg_partials"], idx="tid", wg="WG", form="head", breakInline=true) }}
return sg_partials[0];
{% endif %}
}
// 4 contiguous residual elements (input[idx] + skip[skip_idx] [+ bias]) at vec4
// index `vi`. skip_idx == idx for the normal (non-broadcast) path; for a skip
// that broadcasts across the leading/batch dim uses a folded index.
fn residual_value(idx: u32, skip_idx: u32{% if source.hasBias %}, vi: u32{% endif %}) -> vec4<f32> {
var value = vec4<f32>(input[idx]) + vec4<f32>(skip[skip_idx]);
{% if source.hasBias %}
value = value + vec4<f32>(bias[vi]);
{% endif %}
return value;
}
@compute @workgroup_size(WG, 1, 1)
fn main(
@builtin(workgroup_id) wg_id: vec3<u32>,
@builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
@builtin(subgroup_invocation_id) sg_lane: u32,
@builtin(subgroup_id) sg_id: u32,
@builtin(num_subgroups) num_sg: u32{% endif %}
) {
let row = wg_id.x + wg_id.y * params.rowStride;
if (row >= params.rows) {
return;
}
let tid = lid.x;
let base = row * HIDDEN_V;
{% if source.broadcastSkip %}
// skip broadcasts across the batch dim: fold row into [0, skipRows) so every
// batch reuses the same skip row (skipRows == params.rows ⇒ identity).
let skip_base = (row % params.skipRows) * HIDDEN_V;
{% else %}
let skip_base = base;
{% endif %}
let shift = residual_value(base, skip_base{% if source.hasBias %}, 0u{% endif %}).x;
var acc = vec2<f32>(0.0, 0.0);
for (var i = tid; i < HIDDEN_V; i = i + WG) {
let v = residual_value(base + i, skip_base + i{% if source.hasBias %}, i{% endif %});
let d = v - vec4<f32>(shift);
acc.x = acc.x + d.x + d.y + d.z + d.w;
acc.y = acc.y + dot(d, d);
}
let totals = reduce_pair(acc{% if useSubgroups %}, sg_lane, sg_id, num_sg{% else %}, tid{% endif %});
let mean_d = totals.x / f32(HIDDEN);
let variance = max(totals.y / f32(HIDDEN) - mean_d * mean_d, 0.0);
let row_inv = inverseSqrt(variance + params.epsilon);
let row_mean = shift + mean_d;
for (var i = tid; i < HIDDEN_V; i = i + WG) {
let idx = base + i;
let residual = residual_value(idx, skip_base + i{% if source.hasBias %}, i{% endif %});
{% if source.writeResidualSum %}
input_skip_bias_sum[idx] = {{ source.vecType }}(residual);
{% endif %}
var value = (residual - vec4<f32>(row_mean)) * row_inv * vec4<f32>(gamma[i]);
{% if source.hasBeta %}
value = value + vec4<f32>(beta[i]);
{% endif %}
output[idx] = {{ source.vecType }}(value);
}
}