ai.onnx.InstanceNormalization / build /webgpu /instance-normalization-splitk-partials.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
7eef075 verified
Raw History Blame
4.3 kB
{% 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 %}
/* Split-K partial sums for tensors with few planes and a large spatial extent.
A workgroup-per-plane kernel exposes too little parallelism, so this pass
splits each plane across SPLIT workgroups. Each accumulates a raw sum and
sum-of-squares over its slice. The combine pass produces mean and inverse
standard deviation, and the apply pass normalizes. */
{% set vectorized = vectorized if vectorized is defined else false %}
{% set useSubgroups = useSubgroups if useSubgroups is defined else false %}
{% if useSubgroups %}
enable subgroups;
{% endif %}
{% set LOAD_OPEN = "vec4<f32>(" if usesF16 else "" %}
{% set LOAD_CLOSE = ")" if usesF16 else "" %}
{{ env.wgsl.resourceDeclarations }}
const WG: u32 = {{ workgroupSize }}u;
const SPLIT: u32 = {{ split }}u;
{% if useSubgroups %}
// One slot per possible subgroup avoids assuming any mapping from local
// invocation IDs to subgroup membership.
var<workgroup> subgroup_partials: array<vec2<f32>, WG>;
{% else %}
var<workgroup> red_sum: array<f32, WG>;
var<workgroup> red_sq: array<f32, WG>;
{% endif %}
@compute @workgroup_size(WG, 1, 1)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
@builtin(subgroup_invocation_id) subgroup_lane: u32,
@builtin(subgroup_id) subgroup_id: u32,
@builtin(num_subgroups) num_subgroups: u32{% endif %}) {
let plane = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
if (plane >= params.planes) {
return;
}
let k = wg.z;
let tid = lid.x;
{% if vectorized %}
let spatial = params.spatial / 4u;
{% else %}
let spatial = params.spatial;
{% endif %}
let chunk = (spatial + SPLIT - 1u) / SPLIT;
let start = k * chunk;
var end = start + chunk;
if (end > spatial) { end = spatial; }
let base = plane * spatial;
// Raw second moments cancel: a plane centred on 8192 with unit variance loses
// the variance entirely in E[x^2] - E[x]^2, and the combine's max(.,0) then
// reports zero. Both passes accumulate around the plane's first element, which
// costs one broadcast load and leaves the squared term holding the residual.
{% if vectorized %}
let shift = f32({{ LOAD_OPEN }}input[base]{{ LOAD_CLOSE }}.x);
let shift4 = vec4<f32>(shift);
{% else %}
let shift = f32(input[base]);
{% endif %}
var s = 0.0;
var sq = 0.0;
var i = start + tid;
loop {
if (i >= end) { break; }
{% if vectorized %}
let v = {{ LOAD_OPEN }}input[base + i]{{ LOAD_CLOSE }} - shift4;
s = s + v.x + v.y + v.z + v.w;
sq = sq + dot(v, v);
{% else %}
let v = f32(input[base + i]) - shift;
s = s + v;
sq = sq + v * v;
{% endif %}
i = i + WG;
}
{% if useSubgroups %}
let subgroup_total = vec2<f32>(subgroupAdd(s), subgroupAdd(sq));
if (subgroup_lane == 0u) {
subgroup_partials[subgroup_id] = subgroup_total;
}
workgroupBarrier();
if (tid == 0u) {
var total = vec2<f32>(0.0);
for (var subgroup = 0u; subgroup < num_subgroups; subgroup = subgroup + 1u) {
total = total + subgroup_partials[subgroup];
}
let idx = (plane * SPLIT + k) * 2u;
partials[idx] = total.x;
partials[idx + 1u] = total.y;
}
{% else %}
red_sum[tid] = s;
red_sq[tid] = sq;
workgroupBarrier();
{{ wgsl_tree_fold(["red_sum", "red_sq"], idx="tid", wg="WG", typed=true, form="head", breakInline=true) }}
if (tid == 0u) {
let idx = (plane * SPLIT + k) * 2u;
partials[idx] = red_sum[0];
partials[idx + 1u] = red_sq[0];
}
{% endif %}
}