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