Xenova HF Staff commited on
Commit
7eef075
·
verified ·
1 Parent(s): 8f4239b

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -53,7 +53,7 @@ Some implementation variants require `subgroups`. These are route-specific capab
53
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
54
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
55
  - [`test.json`](build/webgpu/test.json) — correctness cases
56
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
57
  - [`instance-normalization-apply.wgsl.jinja`](build/webgpu/instance-normalization-apply.wgsl.jinja)
58
  - [`instance-normalization-batched-planes-vec4.wgsl.jinja`](build/webgpu/instance-normalization-batched-planes-vec4.wgsl.jinja)
59
  - [`instance-normalization-splitk-combine.wgsl.jinja`](build/webgpu/instance-normalization-splitk-combine.wgsl.jinja)
@@ -63,7 +63,7 @@ Some implementation variants require `subgroups`. These are route-specific capab
63
  ## Use with `@huggingface/kernels`
64
 
65
  ```sh
66
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
67
  ```
68
 
69
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
53
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
54
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
55
  - [`test.json`](build/webgpu/test.json) — correctness cases
56
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
57
  - [`instance-normalization-apply.wgsl.jinja`](build/webgpu/instance-normalization-apply.wgsl.jinja)
58
  - [`instance-normalization-batched-planes-vec4.wgsl.jinja`](build/webgpu/instance-normalization-batched-planes-vec4.wgsl.jinja)
59
  - [`instance-normalization-splitk-combine.wgsl.jinja`](build/webgpu/instance-normalization-splitk-combine.wgsl.jinja)
 
63
  ## Use with `@huggingface/kernels`
64
 
65
  ```sh
66
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
67
  ```
68
 
69
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/bench.json CHANGED
@@ -70,7 +70,7 @@
70
  }
71
  },
72
  {
73
- "name": "alignment_healthy_rank3_2x64x4096_aligned",
74
  "preset": "smoke",
75
  "vars": { "dtype": "float32", "batch": 2, "channels": 64, "spatial": 4096 },
76
  "inputs": {
@@ -108,7 +108,7 @@
108
  }
109
  },
110
  {
111
- "name": "dispatch_healthy_rows60000_under_cap",
112
  "preset": "smoke",
113
  "vars": { "dtype": "float32", "batch": 1, "channels": 60000, "spatial": 64 },
114
  "inputs": {
@@ -205,11 +205,11 @@
205
  }
206
  },
207
  {
208
- "name": "splitk-priority-cliff-c256-256x256",
209
  "preset": "stress",
210
  "provenance": {
211
  "source": "synthetic benchmark",
212
- "notes": "Realistic 16.8M-element feature map that pins the selector boundary between plane_subgroup_vec4 and plane_splitk."
213
  },
214
  "vars": { "dtype": "float32", "batch": 1, "channels": 256, "spatial": 65536 },
215
  "attrs": { "epsilon": 0.00001 },
@@ -229,11 +229,11 @@
229
  }
230
  },
231
  {
232
- "name": "splitk-priority-cliff-c32-512x512",
233
  "preset": "stress",
234
  "provenance": {
235
  "source": "synthetic benchmark",
236
- "notes": "Realistic 8.4M-element high-resolution feature map that pins the selector boundary between plane_subgroup_vec4 and plane_splitk."
237
  },
238
  "vars": { "dtype": "float32", "batch": 1, "channels": 32, "spatial": 262144 },
239
  "attrs": { "epsilon": 0.00001 },
 
70
  }
71
  },
72
  {
73
+ "name": "alignment_control_rank3_2x64x4096_aligned",
74
  "preset": "smoke",
75
  "vars": { "dtype": "float32", "batch": 2, "channels": 64, "spatial": 4096 },
76
  "inputs": {
 
108
  }
109
  },
110
  {
111
+ "name": "dispatch_control_rows60000_under_cap",
112
  "preset": "smoke",
113
  "vars": { "dtype": "float32", "batch": 1, "channels": 60000, "spatial": 64 },
114
  "inputs": {
 
205
  }
206
  },
207
  {
208
+ "name": "splitk-c256-256x256",
209
  "preset": "stress",
210
  "provenance": {
211
  "source": "synthetic benchmark",
212
+ "notes": "A 16.8-million-element feature map measures normalization across large spatial planes."
213
  },
214
  "vars": { "dtype": "float32", "batch": 1, "channels": 256, "spatial": 65536 },
215
  "attrs": { "epsilon": 0.00001 },
 
229
  }
230
  },
231
  {
232
+ "name": "splitk-c32-512x512",
233
  "preset": "stress",
234
  "provenance": {
235
  "source": "synthetic benchmark",
236
+ "notes": "An 8.4-million-element high-resolution feature map measures normalization across large spatial planes."
237
  },
238
  "vars": { "dtype": "float32", "batch": 1, "channels": 32, "spatial": 262144 },
239
  "attrs": { "epsilon": 0.00001 },
build/webgpu/instance-normalization-apply.wgsl.jinja CHANGED
@@ -8,19 +8,21 @@
8
  {% set STORE_CLOSE = ")" if usesF16 else "" %}
9
  {% set CHAN_OPEN = "f32(" if usesF16 else "" %}
10
  {% set CHAN_CLOSE = ")" if usesF16 else "" %}
 
 
 
 
 
 
 
 
11
  {{ env.wgsl.resourceDeclarations }}
12
 
13
  const WG: u32 = {{ applyWorkgroupSize }}u;
14
 
15
  @compute @workgroup_size(WG, 1, 1)
16
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
17
- {% if not vectorized %}
18
- // 2D-folded flat index: gid.y carries the high bits after dispatch folding.
19
- {% endif %}
20
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
21
- if (index >= params.count) {
22
- return;
23
- }
24
  {% if vectorized %}
25
  // The vectorized path requires each plane to contain a multiple of four
26
  // values, so a packed load/store never crosses an instance boundary.
 
8
  {% set STORE_CLOSE = ")" if usesF16 else "" %}
9
  {% set CHAN_OPEN = "f32(" if usesF16 else "" %}
10
  {% set CHAN_CLOSE = ")" if usesF16 else "" %}
11
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
12
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
13
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
14
+ // per-axis workgroup fold width.
15
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
16
+ if ({{ name }} >= {{ bound }}) {
17
+ return;
18
+ }{% endmacro %}
19
  {{ env.wgsl.resourceDeclarations }}
20
 
21
  const WG: u32 = {{ applyWorkgroupSize }}u;
22
 
23
  @compute @workgroup_size(WG, 1, 1)
24
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
25
+ {{ flat_index_2d("WG", "index") }}
 
 
 
 
 
 
26
  {% if vectorized %}
27
  // The vectorized path requires each plane to contain a multiple of four
28
  // values, so a packed load/store never crosses an instance boundary.
build/webgpu/instance-normalization-splitk-combine.wgsl.jinja CHANGED
@@ -2,6 +2,14 @@
2
  // standard deviation. One thread handles each plane. The partials are centred on
3
  // the plane's first element, so E[y^2] - E[y]^2 keeps the variance a raw second
4
  // moment would cancel away; max(value, 0) guards against negative rounding residue.
 
 
 
 
 
 
 
 
5
  {{ env.wgsl.resourceDeclarations }}
6
 
7
  const SPLIT: u32 = {{ split }}u;
@@ -9,10 +17,7 @@ const COMBINE_WG: u32 = {{ combineWorkgroupSize }}u;
9
 
10
  @compute @workgroup_size(COMBINE_WG, 1, 1)
11
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
- let plane = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * COMBINE_WG;
13
- if (plane >= params.planes) {
14
- return;
15
- }
16
  var total = 0.0;
17
  var total_sq = 0.0;
18
  let b = plane * SPLIT;
 
2
  // standard deviation. One thread handles each plane. The partials are centred on
3
  // the plane's first element, so E[y^2] - E[y]^2 keeps the variance a raw second
4
  // moment would cancel away; max(value, 0) guards against negative rounding residue.
5
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
6
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
7
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
8
+ // per-axis workgroup fold width.
9
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
10
+ if ({{ name }} >= {{ bound }}) {
11
+ return;
12
+ }{% endmacro %}
13
  {{ env.wgsl.resourceDeclarations }}
14
 
15
  const SPLIT: u32 = {{ split }}u;
 
17
 
18
  @compute @workgroup_size(COMBINE_WG, 1, 1)
19
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
20
+ {{ flat_index_2d("COMBINE_WG", "plane", "params.planes") }}
 
 
 
21
  var total = 0.0;
22
  var total_sq = 0.0;
23
  let b = plane * SPLIT;
build/webgpu/instance-normalization-splitk-partials.wgsl.jinja CHANGED
@@ -1,49 +1,19 @@
1
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
2
- {% if op == "max" %}
3
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
4
- {%- else %}
5
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
6
- {%- endif %}
7
- {% endmacro %}
8
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
9
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
10
  loop {
11
- {% if form == "head" %}
12
- {% if breakInline %}
13
  if ({{ svar }} == 0u) { break; }
14
- {% else %}
15
- if ({{ svar }} == 0u) {
16
- break;
17
- }
18
- {% endif %}
19
- {% endif %}
20
- {% if bodyInline %}
21
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
22
- {% else %}
23
  if ({{ idx }} < {{ svar }}) {
24
  {% for a in arrays %}
25
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
26
  {% endfor %}
27
  }
28
- {% endif %}
29
- {% if form == "head" %}
30
- {% if barrierFirst %}
31
- workgroupBarrier();
32
- {{ svar }} = {{ svar }} / 2u;
33
- {% else %}
34
  {{ svar }} = {{ svar }} / 2u;
35
  workgroupBarrier();
36
- {% endif %}
37
- {% else %}
38
- workgroupBarrier();
39
- if ({{ svar }} == 1u) {
40
- break;
41
- }
42
- {{ svar }} = {{ svar }} / 2u;
43
- {% endif %}
44
- }
45
- {%- endmacro %}
46
-
47
  /* Split-K partial sums for tensors with few planes and a large spatial extent.
48
  A workgroup-per-plane kernel exposes too little parallelism, so this pass
49
  splits each plane across SPLIT workgroups. Each accumulates a raw sum and
 
1
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
2
+ {% if op == "max" or op == "min" %}
3
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
4
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
5
+ {% 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) %}
 
 
 
6
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
7
  loop {
 
 
8
  if ({{ svar }} == 0u) { break; }
 
 
 
 
 
 
 
 
 
9
  if ({{ idx }} < {{ svar }}) {
10
  {% for a in arrays %}
11
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
12
  {% endfor %}
13
  }
 
 
 
 
 
 
14
  {{ svar }} = {{ svar }} / 2u;
15
  workgroupBarrier();
16
+ }{% endmacro %}
 
 
 
 
 
 
 
 
 
 
17
  /* Split-K partial sums for tensors with few planes and a large spatial extent.
18
  A workgroup-per-plane kernel exposes too little parallelism, so this pass
19
  splits each plane across SPLIT workgroups. Each accumulates a raw sum and
build/webgpu/manifest.json CHANGED
@@ -48,28 +48,26 @@
48
  "splitStatsPreferred": "splitStatsCovered and instancePlanes < normSubgroupMax"
49
  },
50
  "bindings": {
51
- "x": { "arg": "input", "buffer": "read-only-storage", "elementType": "$ioElement" },
52
- "scale": { "buffer": "read-only-storage", "elementType": "$T" },
53
- "bias": { "arg": "b", "buffer": "read-only-storage", "elementType": "$T" },
54
- "y": { "arg": "output", "buffer": "storage", "elementType": "$ioElement" },
55
- "input": { "buffer": "read-only-storage", "elementType": "$splitInputElement" },
56
- "input_2": { "name": "input", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
57
- "stats_2": { "name": "stats", "buffer": "read-only-storage", "elementType": "f32" },
58
- "output": { "buffer": "storage", "elementType": "$vectorScalar" },
59
- "params_5": {
60
  "name": "params",
61
- "buffer": "uniform",
62
  "struct": [
63
  { "name": "count", "type": "u32", "value": "numel(shapes.output) / 4" },
64
  { "name": "channels", "type": "u32", "value": "dim(shapes.input, 1)" },
65
  { "name": "spatial", "type": "u32", "value": "instanceSpatial" }
66
  ]
67
  },
68
- "input_3": { "name": "input", "buffer": "read-only-storage", "elementType": "$T" },
69
- "output_2": { "name": "output", "buffer": "storage", "elementType": "$T" },
70
- "params_6": {
71
  "name": "params",
72
- "buffer": "uniform",
73
  "struct": [
74
  { "name": "count", "type": "u32", "value": "numel(shapes.output)" },
75
  { "name": "channels", "type": "u32", "value": "dim(shapes.input, 1)" },
@@ -84,7 +82,6 @@
84
  "when": ["instanceRowCovered", "instanceSpatial % 4 == 0", "instanceSpatial >= 4", "instancePlanes >= normWorkgroupCap", "instanceBatchedVec4PlanesPerWorkgroup >= tunables.BATCHED_MIN_PLANES_PER_WORKGROUP", "instanceBatchedVec4StorageBytes <= device.limits.maxComputeWorkgroupStorageSize"],
85
  "demoteWhen": ["reportedNonWave32Adapter and instancePlanes <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
86
  "derive": {
87
- "usesF16": "dtypes.T == \"f16\"",
88
  "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"",
89
  "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
90
  "hidden": "instanceSpatial",
@@ -123,17 +120,15 @@
123
  "priority": 110,
124
  "when": ["instanceRowCovered", "inner(shapes.input, 1) % 4 == 0", "instanceVec4SubgroupEfficient"],
125
  "requires": { "features": [] },
126
- "derive": { "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
127
  "passes": [
128
  {
129
  "id": "main",
130
- "name": "InstanceNormalization.plane_subgroup_vec4",
131
  "shader": "norm-row-stats.wgsl.jinja",
132
  "derive": {
133
- "modeSpec": "\"instance\"",
134
  "vec4": true,
135
  "scalar": "dtypes.T",
136
- "usesF16Spec": "dtypes.T == \"f16\"",
137
  "hidden": "instanceSpatial",
138
  "wg": "instanceVec4Workgroup",
139
  "epsilon": "attrs.epsilon",
@@ -159,8 +154,7 @@
159
  ]
160
  }
161
  ],
162
- "dispatch": { "x": "min(instancePlanes, 65535)", "y": "ceilDiv(instancePlanes, 65535)", "z": 1 },
163
- "subgroupCollectivesWidth": "portable"
164
  }
165
  ]
166
  },
@@ -173,14 +167,12 @@
173
  "passes": [
174
  {
175
  "id": "main",
176
- "name": "InstanceNormalization.plane_subgroup_vec4_scalar_io",
177
  "shader": "norm-row-stats.wgsl.jinja",
178
  "derive": {
179
- "modeSpec": "\"instance\"",
180
  "vec4": true,
181
  "scalarIo": true,
182
  "scalar": "dtypes.T",
183
- "usesF16Spec": false,
184
  "hidden": "instanceSpatial",
185
  "wg": "instanceVec4Workgroup",
186
  "epsilon": "attrs.epsilon",
@@ -206,8 +198,7 @@
206
  ]
207
  }
208
  ],
209
- "dispatch": { "x": "min(instancePlanes, 65535)", "y": "ceilDiv(instancePlanes, 65535)", "z": 1 },
210
- "subgroupCollectivesWidth": "portable"
211
  }
212
  ]
213
  },
@@ -220,13 +211,11 @@
220
  "passes": [
221
  {
222
  "id": "main",
223
- "name": "InstanceNormalization.plane_subgroup",
224
  "shader": "norm-row-stats.wgsl.jinja",
225
  "derive": {
226
- "modeSpec": "\"instance\"",
227
  "vec4": false,
228
  "scalar": "dtypes.T",
229
- "usesF16Spec": "dtypes.T == \"f16\"",
230
  "hidden": "instanceSpatial",
231
  "wg": "instanceScalarWorkgroup",
232
  "epsilon": "attrs.epsilon",
@@ -250,8 +239,7 @@
250
  ]
251
  }
252
  ],
253
- "dispatch": { "x": "min(instancePlanes, 65535)", "y": "ceilDiv(instancePlanes, 65535)", "z": 1 },
254
- "subgroupCollectivesWidth": "portable"
255
  }
256
  ]
257
  },
@@ -281,7 +269,7 @@
281
  "shader": "instance-normalization-splitk-partials.wgsl.jinja",
282
  "bindings": [
283
  "input",
284
- { "name": "partials", "buffer": "storage", "elementType": "f32" },
285
  {
286
  "name": "params",
287
  "struct": [
@@ -294,8 +282,7 @@
294
  "x": "min(instancePlanes, DISPATCH_FOLD_WIDTH)",
295
  "y": "ceilDiv(instancePlanes, DISPATCH_FOLD_WIDTH)",
296
  "z": "instanceSplitCount"
297
- },
298
- "subgroupCollectivesWidth": "portable"
299
  },
300
  {
301
  "id": "combine",
@@ -305,7 +292,7 @@
305
  "bindings": [
306
  "input",
307
  { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
308
- { "name": "stats", "buffer": "storage", "elementType": "f32" },
309
  {
310
  "name": "params",
311
  "struct": [
@@ -325,7 +312,7 @@
325
  "id": "apply",
326
  "name": "InstanceNormalization.ApplyVec4",
327
  "shader": "instance-normalization-apply.wgsl.jinja",
328
- "bindings": ["input_2", "stats_2", "scale", "bias", "output", "params_5"],
329
  "dispatch": {
330
  "x": "min(ceilDiv((numel(shapes.output) / 4), (applyWorkgroupSize)), 65535)",
331
  "y": "ceilDiv(ceilDiv((numel(shapes.output) / 4), (applyWorkgroupSize)), 65535)",
@@ -358,7 +345,7 @@
358
  "shader": "instance-normalization-splitk-partials.wgsl.jinja",
359
  "bindings": [
360
  "input",
361
- { "name": "partials", "buffer": "storage", "elementType": "f32" },
362
  {
363
  "name": "params",
364
  "struct": [
@@ -381,7 +368,7 @@
381
  "bindings": [
382
  "input",
383
  { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
384
- { "name": "stats", "buffer": "storage", "elementType": "f32" },
385
  {
386
  "name": "params",
387
  "struct": [
@@ -401,7 +388,7 @@
401
  "id": "apply",
402
  "name": "InstanceNormalization.Apply",
403
  "shader": "instance-normalization-apply.wgsl.jinja",
404
- "bindings": ["input_3", "stats_2", "scale", "bias", "output_2", "params_6"],
405
  "dispatch": {
406
  "x": "min(ceilDiv((numel(shapes.output)), (applyWorkgroupSize)), 65535)",
407
  "y": "ceilDiv(ceilDiv((numel(shapes.output)), (applyWorkgroupSize)), 65535)",
 
48
  "splitStatsPreferred": "splitStatsCovered and instancePlanes < normSubgroupMax"
49
  },
50
  "bindings": {
51
+ "x": { "arg": "input", "elementType": "$ioElement" },
52
+ "scale": { "elementType": "$T" },
53
+ "bias": { "arg": "b", "elementType": "$T" },
54
+ "y": { "arg": "output", "elementType": "$ioElement" },
55
+ "input": { "elementType": "$splitInputElement" },
56
+ "input_apply": { "name": "input", "elementType": "$vectorScalar" },
57
+ "stats_f32": { "name": "stats", "buffer": "read-only-storage", "elementType": "f32" },
58
+ "output": { "elementType": "$vectorScalar" },
59
+ "params_apply": {
60
  "name": "params",
 
61
  "struct": [
62
  { "name": "count", "type": "u32", "value": "numel(shapes.output) / 4" },
63
  { "name": "channels", "type": "u32", "value": "dim(shapes.input, 1)" },
64
  { "name": "spatial", "type": "u32", "value": "instanceSpatial" }
65
  ]
66
  },
67
+ "input_t": { "name": "input", "elementType": "$T" },
68
+ "output_t": { "name": "output", "elementType": "$T" },
69
+ "params__uniform": {
70
  "name": "params",
 
71
  "struct": [
72
  { "name": "count", "type": "u32", "value": "numel(shapes.output)" },
73
  { "name": "channels", "type": "u32", "value": "dim(shapes.input, 1)" },
 
82
  "when": ["instanceRowCovered", "instanceSpatial % 4 == 0", "instanceSpatial >= 4", "instancePlanes >= normWorkgroupCap", "instanceBatchedVec4PlanesPerWorkgroup >= tunables.BATCHED_MIN_PLANES_PER_WORKGROUP", "instanceBatchedVec4StorageBytes <= device.limits.maxComputeWorkgroupStorageSize"],
83
  "demoteWhen": ["reportedNonWave32Adapter and instancePlanes <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
84
  "derive": {
 
85
  "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"",
86
  "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
87
  "hidden": "instanceSpatial",
 
120
  "priority": 110,
121
  "when": ["instanceRowCovered", "inner(shapes.input, 1) % 4 == 0", "instanceVec4SubgroupEfficient"],
122
  "requires": { "features": [] },
123
+ "derive": { "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" },
124
  "passes": [
125
  {
126
  "id": "main",
127
+ "name": "InstanceNormalization.PlaneSubgroupVec4",
128
  "shader": "norm-row-stats.wgsl.jinja",
129
  "derive": {
 
130
  "vec4": true,
131
  "scalar": "dtypes.T",
 
132
  "hidden": "instanceSpatial",
133
  "wg": "instanceVec4Workgroup",
134
  "epsilon": "attrs.epsilon",
 
154
  ]
155
  }
156
  ],
157
+ "dispatch": { "x": "min(instancePlanes, 65535)", "y": "ceilDiv(instancePlanes, 65535)", "z": 1 }
 
158
  }
159
  ]
160
  },
 
167
  "passes": [
168
  {
169
  "id": "main",
170
+ "name": "InstanceNormalization.PlaneSubgroupVec4ScalarIo",
171
  "shader": "norm-row-stats.wgsl.jinja",
172
  "derive": {
 
173
  "vec4": true,
174
  "scalarIo": true,
175
  "scalar": "dtypes.T",
 
176
  "hidden": "instanceSpatial",
177
  "wg": "instanceVec4Workgroup",
178
  "epsilon": "attrs.epsilon",
 
198
  ]
199
  }
200
  ],
201
+ "dispatch": { "x": "min(instancePlanes, 65535)", "y": "ceilDiv(instancePlanes, 65535)", "z": 1 }
 
202
  }
203
  ]
204
  },
 
211
  "passes": [
212
  {
213
  "id": "main",
214
+ "name": "InstanceNormalization.PlaneSubgroup",
215
  "shader": "norm-row-stats.wgsl.jinja",
216
  "derive": {
 
217
  "vec4": false,
218
  "scalar": "dtypes.T",
 
219
  "hidden": "instanceSpatial",
220
  "wg": "instanceScalarWorkgroup",
221
  "epsilon": "attrs.epsilon",
 
239
  ]
240
  }
241
  ],
242
+ "dispatch": { "x": "min(instancePlanes, 65535)", "y": "ceilDiv(instancePlanes, 65535)", "z": 1 }
 
243
  }
244
  ]
245
  },
 
269
  "shader": "instance-normalization-splitk-partials.wgsl.jinja",
270
  "bindings": [
271
  "input",
272
+ { "name": "partials", "elementType": "f32" },
273
  {
274
  "name": "params",
275
  "struct": [
 
282
  "x": "min(instancePlanes, DISPATCH_FOLD_WIDTH)",
283
  "y": "ceilDiv(instancePlanes, DISPATCH_FOLD_WIDTH)",
284
  "z": "instanceSplitCount"
285
+ }
 
286
  },
287
  {
288
  "id": "combine",
 
292
  "bindings": [
293
  "input",
294
  { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
295
+ { "name": "stats", "elementType": "f32" },
296
  {
297
  "name": "params",
298
  "struct": [
 
312
  "id": "apply",
313
  "name": "InstanceNormalization.ApplyVec4",
314
  "shader": "instance-normalization-apply.wgsl.jinja",
315
+ "bindings": ["input_apply", "stats_f32", "scale", "bias", "output", "params_apply"],
316
  "dispatch": {
317
  "x": "min(ceilDiv((numel(shapes.output) / 4), (applyWorkgroupSize)), 65535)",
318
  "y": "ceilDiv(ceilDiv((numel(shapes.output) / 4), (applyWorkgroupSize)), 65535)",
 
345
  "shader": "instance-normalization-splitk-partials.wgsl.jinja",
346
  "bindings": [
347
  "input",
348
+ { "name": "partials", "elementType": "f32" },
349
  {
350
  "name": "params",
351
  "struct": [
 
368
  "bindings": [
369
  "input",
370
  { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
371
+ { "name": "stats", "elementType": "f32" },
372
  {
373
  "name": "params",
374
  "struct": [
 
388
  "id": "apply",
389
  "name": "InstanceNormalization.Apply",
390
  "shader": "instance-normalization-apply.wgsl.jinja",
391
+ "bindings": ["input_t", "stats_f32", "scale", "bias", "output_t", "params__uniform"],
392
  "dispatch": {
393
  "x": "min(ceilDiv((numel(shapes.output)), (applyWorkgroupSize)), 65535)",
394
  "y": "ceilDiv(ceilDiv((numel(shapes.output)), (applyWorkgroupSize)), 65535)",
build/webgpu/metadata.json CHANGED
@@ -1,25 +1,25 @@
1
  {
2
  "name": "ai.onnx.InstanceNormalization",
3
- "id": "_ai_onnx_instancenormalization_webgpu_16b3576",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "HBqAImIaa6N6gGqA53oRQ/ZvaCAFNJPKltDsBnjF5wQ=",
11
- "instance-normalization-apply.wgsl.jinja": "ss27/JnTr5XXoUd0lj8yyYcXfYImAO5jPB9BdkneLwU=",
12
  "instance-normalization-batched-planes-vec4.wgsl.jinja": "cLkDhQOaM/T+im43mRMyLa+kEoIHfm8i31IdkiyeMcI=",
13
- "instance-normalization-splitk-combine.wgsl.jinja": "z2uqoUYCw3foyDxmSSD9i8fjH3Z6HFpzVhBlNCV11BM=",
14
- "instance-normalization-splitk-partials.wgsl.jinja": "SymImIFJ0cyKgoJZtsuTtPAEy2m8CtCKWn8ylnuBQJE=",
15
- "manifest.json": "Eq6DFSvT5bXhoGN7gsrUL106yaDyqRNzf+r4t2e/Nh0=",
16
- "norm-row-stats.wgsl.jinja": "pIQ85rOY7wxOVoImYTZF8yUjKpdvCKAH4h2EBmD1Kvc=",
17
- "test.json": "AIj6JyZjhnniT+w2T4DZkIjVYrltvKmjY6cFQcNm+kY="
18
  }
19
  },
20
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
21
  "webgpu": {
22
- "manifestSpec": "2.0",
23
  "variants": {
24
  "plane_batched_vec4": ["instance-normalization-batched-planes-vec4.wgsl.jinja"],
25
  "plane_subgroup_vec4": ["norm-row-stats.wgsl.jinja"],
 
1
  {
2
  "name": "ai.onnx.InstanceNormalization",
3
+ "id": "_ai_onnx_instancenormalization_webgpu_c05b966",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "9QDOAZ+rJkbSQmr6aXgD0P9GVbRhHp7Dv8hb3qshhE4=",
11
+ "instance-normalization-apply.wgsl.jinja": "0RQER6S3FawT3PiI85Q1q+uei0OuCB3XcNwQ7+KC4zw=",
12
  "instance-normalization-batched-planes-vec4.wgsl.jinja": "cLkDhQOaM/T+im43mRMyLa+kEoIHfm8i31IdkiyeMcI=",
13
+ "instance-normalization-splitk-combine.wgsl.jinja": "W+2LgCldnWkp2Lw0pl9fGRX3YNgM0Y4rIHIuWSDYCds=",
14
+ "instance-normalization-splitk-partials.wgsl.jinja": "WeKaPBbMMmRWbiW09cUT00LOhQITXbKXmHiZgRBzFSQ=",
15
+ "manifest.json": "Y69yrN2/0rhubLl2fPJ7WSXnZ6HbxwUI0KZX2lmY6lg=",
16
+ "norm-row-stats.wgsl.jinja": "lOpzwxHA02sP4YiHrqrV8O75Bfj2SFpGz18yLU61w6g=",
17
+ "test.json": "371aSdeBp/Pdxtdo8XUWTLJFBaUJVatpfPDqLNU9DbM="
18
  }
19
  },
20
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
21
  "webgpu": {
22
+ "manifestSpec": "2.1",
23
  "variants": {
24
  "plane_batched_vec4": ["instance-normalization-batched-planes-vec4.wgsl.jinja"],
25
  "plane_subgroup_vec4": ["norm-row-stats.wgsl.jinja"],
build/webgpu/norm-row-stats.wgsl.jinja CHANGED
@@ -1,15 +1,5 @@
1
- {% if usesF16Spec %}
2
- enable f16;
3
- {% endif %}
4
- {% set combineSubgroups = combineSubgroups %}
5
  {% set scalarIo = scalarIo if scalarIo is defined else false %}
6
- {% set packedBf16Embedding = packedBf16Embedding if packedBf16Embedding is defined else false %}
7
- {% set rmsChainNorm = rmsChainNorm if rmsChainNorm is defined else false %}
8
- {% set hiddenPairs = hiddenPairs | default(0) %}
9
- {% set numRows = numRows | default(0) %}
10
- {% set epsilon = epsilon | default("0.0") %}
11
- {% set epsilon2 = epsilon2 | default("0.0") %}
12
- {% set channels = channels | default(0) %}
13
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
14
  if combineSubgroups else ", tid: u32" %}
15
  {% set reduceThreadArguments = ", sg_lane, sg_id, num_sg"
@@ -33,49 +23,9 @@ const HIDDEN: u32 = {{ hidden }}u;
33
  {% if vec4 %}
34
  const HIDDEN_V: u32 = {{ hiddenVec }}u;
35
  {% endif %}
36
- {% if packedBf16Embedding %}
37
- const HIDDEN_PAIRS: u32 = {{ hiddenPairs }}u;
38
- const NUM_ROWS: u32 = {{ numRows }}u;
39
- {% endif %}
40
  const WG: u32 = {{ wg }}u;
41
  const EPSILON: f32 = {{ epsilon }};
42
- {% if rmsChainNorm %}
43
- const EPSILON2: f32 = {{ epsilon2 }};
44
- {% endif %}
45
  const CHANNELS: u32 = {{ channels }}u;
46
-
47
- {% if packedBf16Embedding %}
48
- {% if vec4 %}
49
- fn unpack_bf16_pair(word: u32) -> vec2<f32> {
50
- let bits = vec2<u32>(word & 0xffffu, word >> 16u);
51
- return bitcast<vec2<f32>>(bits << vec2<u32>(16u));
52
- }
53
- {% endif %}
54
-
55
- {% if not vec4 %}
56
- fn embedding_scalar(source_row: u32, hidden: u32) -> f32 {
57
- if (source_row >= NUM_ROWS) {
58
- return 0.0;
59
- }
60
- let word = x[source_row * HIDDEN_PAIRS + (hidden >> 1u)];
61
- let bits = select(word & 0xffffu, word >> 16u, (hidden & 1u) != 0u);
62
- return bitcast<f32>(bits << 16u);
63
- }
64
- {% endif %}
65
-
66
- {% if vec4 %}
67
- fn embedding_vec4(source_row: u32, hidden_vec: u32) -> vec4<f32> {
68
- if (source_row >= NUM_ROWS) {
69
- return vec4<f32>(0.0);
70
- }
71
- let base = source_row * HIDDEN_PAIRS + hidden_vec * 2u;
72
- let low = unpack_bf16_pair(x[base]);
73
- let high = unpack_bf16_pair(x[base + 1u]);
74
- return vec4<f32>(low, high);
75
- }
76
- {% endif %}
77
- {% endif %}
78
-
79
  {% if vec4 and scalarIo %}
80
  fn load_vec4(index: u32) -> vec4<f32> {
81
  return vec4<f32>(x[index], x[index + 1u], x[index + 2u], x[index + 3u]);
@@ -139,14 +89,7 @@ fn main(
139
  return;
140
  }
141
  let tid = lid.x;
142
- {% if packedBf16Embedding %}
143
- let source_row = indices[row];
144
- {% if vec4 %}
145
- let base = row * HIDDEN_V;
146
- {% else %}
147
- let base = row * HIDDEN;
148
- {% endif %}
149
- {% elif vec4 and not scalarIo %}
150
  let base = row * HIDDEN_V;
151
  {% else %}
152
  let base = row * HIDDEN;
@@ -165,26 +108,18 @@ fn main(
165
  var acc = vec2<f32>(0.0, 0.0);
166
  {% if vec4 %}
167
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
168
- {% if packedBf16Embedding %}
169
- let v = embedding_vec4(source_row, i);
170
- embedding_out[base + i] = v;
171
- {% elif scalarIo %}
172
  let v = load_vec4(base + i * 4u);
173
  {% else %}
174
- let v = vec4<f32>(x[base + i]);
175
  {% endif %}
176
- let d = v - vec4<f32>(shift);
177
  acc.x = acc.x + d.x + d.y + d.z + d.w;
178
  acc.y = acc.y + dot(d, d);
179
  }
180
  {% else %}
181
  for (var i = tid; i < HIDDEN; i = i + WG) {
182
- {% if packedBf16Embedding %}
183
- let v = embedding_scalar(source_row, i);
184
- embedding_out[base + i] = v;
185
- {% else %}
186
  let v = f32(x[base + i]);
187
- {% endif %}
188
  let d = v - shift;
189
  acc.x = acc.x + d;
190
  acc.y = acc.y + d * d;
@@ -201,20 +136,14 @@ fn main(
201
  let ch_scale = f32(scale[c]);
202
  let ch_bias = f32(bias[c]);
203
 
204
- {% if rmsChainNorm %}
205
- var acc2 = 0.0;
206
- {% endif %}
207
  {% if vec4 %}
208
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
209
- {% if packedBf16Embedding %}
210
- let idx = base + i;
211
- let v = embedding_vec4(source_row, i);
212
- {% elif scalarIo %}
213
  let idx = base + i * 4u;
214
  let v = load_vec4(idx);
215
  {% else %}
216
  let idx = base + i;
217
- let v = vec4<f32>(x[idx]);
218
  {% endif %}
219
  {% if scalarIo %}
220
  let value = (v - vec4<f32>(row_mean)) * inv * vec4<f32>(ch_scale) + vec4<f32>(ch_bias);
@@ -226,29 +155,10 @@ fn main(
226
  y[idx] = {{ vecType }}((v - vec4<f32>(row_mean)) * inv * vec4<f32>(ch_scale) + vec4<f32>(ch_bias));
227
  {% endif %}
228
  }
229
- {% if rmsChainNorm %}
230
-
231
- // The chained second norm reads the residual row this loop just stored. This
232
- // barrier completes those stores and any preceding shared-scratch use before
233
- // the next reduction reuses its scratch; each lane then re-reads only the
234
- // elements it wrote itself.
235
- workgroupBarrier();
236
- let total2 = reduce_scalar(acc2{{ reduceThreadArguments }});
237
- let inv2 = inverseSqrt(total2 / f32(HIDDEN) + EPSILON2);
238
- for (var i = tid; i < HIDDEN_V; i = i + WG) {
239
- let idx = base + i;
240
- let hv = vec4<f32>(y[idx]);
241
- normed2[idx] = {{ vecType }}(hv * inv2 * vec4<f32>(scale2[i]));
242
- }
243
- {% endif %}
244
  {% else %}
245
  for (var i = tid; i < HIDDEN; i = i + WG) {
246
  let idx = base + i;
247
- {% if packedBf16Embedding %}
248
- let v = embedding_scalar(source_row, i);
249
- {% else %}
250
  let v = f32(x[idx]);
251
- {% endif %}
252
  y[idx] = {{ scalar }}((v - row_mean) * inv * ch_scale + ch_bias);
253
  }
254
  {% endif %}
 
 
 
 
 
1
  {% set scalarIo = scalarIo if scalarIo is defined else false %}
2
+ {% set packedF32 = "vec4<f32>" %}
 
 
 
 
 
 
3
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
4
  if combineSubgroups else ", tid: u32" %}
5
  {% set reduceThreadArguments = ", sg_lane, sg_id, num_sg"
 
23
  {% if vec4 %}
24
  const HIDDEN_V: u32 = {{ hiddenVec }}u;
25
  {% endif %}
 
 
 
 
26
  const WG: u32 = {{ wg }}u;
27
  const EPSILON: f32 = {{ epsilon }};
 
 
 
28
  const CHANNELS: u32 = {{ channels }}u;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
  {% if vec4 and scalarIo %}
30
  fn load_vec4(index: u32) -> vec4<f32> {
31
  return vec4<f32>(x[index], x[index + 1u], x[index + 2u], x[index + 3u]);
 
89
  return;
90
  }
91
  let tid = lid.x;
92
+ {% if vec4 and not scalarIo %}
 
 
 
 
 
 
 
93
  let base = row * HIDDEN_V;
94
  {% else %}
95
  let base = row * HIDDEN;
 
108
  var acc = vec2<f32>(0.0, 0.0);
109
  {% if vec4 %}
110
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
111
+ {% if scalarIo %}
 
 
 
112
  let v = load_vec4(base + i * 4u);
113
  {% else %}
114
+ let v = {{ packedF32 }}(x[base + i]);
115
  {% endif %}
116
+ let d = v - {{ packedF32 }}(shift);
117
  acc.x = acc.x + d.x + d.y + d.z + d.w;
118
  acc.y = acc.y + dot(d, d);
119
  }
120
  {% else %}
121
  for (var i = tid; i < HIDDEN; i = i + WG) {
 
 
 
 
122
  let v = f32(x[base + i]);
 
123
  let d = v - shift;
124
  acc.x = acc.x + d;
125
  acc.y = acc.y + d * d;
 
136
  let ch_scale = f32(scale[c]);
137
  let ch_bias = f32(bias[c]);
138
 
 
 
 
139
  {% if vec4 %}
140
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
141
+ {% if scalarIo %}
 
 
 
142
  let idx = base + i * 4u;
143
  let v = load_vec4(idx);
144
  {% else %}
145
  let idx = base + i;
146
+ let v = {{ packedF32 }}(x[idx]);
147
  {% endif %}
148
  {% if scalarIo %}
149
  let value = (v - vec4<f32>(row_mean)) * inv * vec4<f32>(ch_scale) + vec4<f32>(ch_bias);
 
155
  y[idx] = {{ vecType }}((v - vec4<f32>(row_mean)) * inv * vec4<f32>(ch_scale) + vec4<f32>(ch_bias));
156
  {% endif %}
157
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
158
  {% else %}
159
  for (var i = tid; i < HIDDEN; i = i + WG) {
160
  let idx = base + i;
 
 
 
161
  let v = f32(x[idx]);
 
162
  y[idx] = {{ scalar }}((v - row_mean) * inv * ch_scale + ch_bias);
163
  }
164
  {% endif %}
build/webgpu/test.json CHANGED
@@ -88,9 +88,7 @@
88
  },
89
  {
90
  "name": "f32_batched_planes_1x257x64",
91
- "provenance": {
92
- "notes": "Compact correctness lock for the feature-independent lane-cohort plane batching used by the 70,000-row dispatch-cliff benchmark."
93
- },
94
  "attrs": { "epsilon": 0.00001 },
95
  "inputs": {
96
  "input": {
 
88
  },
89
  {
90
  "name": "f32_batched_planes_1x257x64",
91
+ "provenance": { "notes": "A compact many-plane input checks independent channel normalization." },
 
 
92
  "attrs": { "epsilon": 0.00001 },
93
  "inputs": {
94
  "input": {