File size: 4,302 Bytes
0af9165 7eef075 0af9165 7eef075 0af9165 8f4239b 0af9165 8f4239b 0af9165 8f4239b 0af9165 8f4239b 0af9165 8f4239b 0af9165 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 | {% 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 %}
}
|