ai.onnx.InstanceNormalization / build /webgpu /instance-normalization-splitk-partials.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 2e7068faf55e
0af9165 verified
Raw
History Blame
4.31 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 %}
/* 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 usesF16 %}
enable f16;
{% endif %}
{% 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>,
@builtin(num_workgroups) nwg: 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 * nwg.x;
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;
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 }};
s = s + v.x + v.y + v.z + v.w;
sq = sq + dot(v, v);
{% else %}
let v = f32(input[base + i]);
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 %}
}