{% 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(" 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 subgroup_partials: array, WG>; {% else %} var red_sum: array; var red_sq: array; {% endif %} @compute @workgroup_size(WG, 1, 1) fn main(@builtin(workgroup_id) wg: vec3, @builtin(local_invocation_id) lid: vec3, @builtin(num_workgroups) nwg: vec3{% 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(subgroupAdd(s), subgroupAdd(sq)); if (subgroup_lane == 0u) { subgroup_partials[subgroup_id] = subgroup_total; } workgroupBarrier(); if (tid == 0u) { var total = vec2(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 %} }