File size: 4,314 Bytes
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 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | {% 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 %}
}
|