| {% 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 %} |
| } |
| |