Xenova HF Staff commited on
Commit
3aa6d13
·
verified ·
1 Parent(s): 0f5881e

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -46,7 +46,7 @@ Attributes and default values (overridable per request):
46
 
47
  | Variable | Allowed dtypes |
48
  | --- | --- |
49
- | `T` | `float32`, `float16`, `uint32`, `int32`, `int16`, `uint8`, `int8`, `bool` |
50
  | `S` | `int64` |
51
 
52
  ## Files
@@ -54,7 +54,7 @@ Attributes and default values (overridable per request):
54
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
55
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
56
  - [`test.json`](build/webgpu/test.json) — correctness cases
57
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
58
  - [`datamove-flat-copy.wgsl.jinja`](build/webgpu/datamove-flat-copy.wgsl.jinja)
59
  - [`datamove-split-block.wgsl.jinja`](build/webgpu/datamove-split-block.wgsl.jinja)
60
  - [`split-n.wgsl.jinja`](build/webgpu/split-n.wgsl.jinja)
@@ -62,7 +62,7 @@ Attributes and default values (overridable per request):
62
  ## Use with `@huggingface/kernels`
63
 
64
  ```sh
65
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
66
  ```
67
 
68
  Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
 
46
 
47
  | Variable | Allowed dtypes |
48
  | --- | --- |
49
+ | `T` | `float32`, `float16`, `uint32`, `int32`, `int16`, `uint8`, `int8`, `bool`, `int64` |
50
  | `S` | `int64` |
51
 
52
  ## Files
 
54
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
55
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
56
  - [`test.json`](build/webgpu/test.json) — correctness cases
57
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
58
  - [`datamove-flat-copy.wgsl.jinja`](build/webgpu/datamove-flat-copy.wgsl.jinja)
59
  - [`datamove-split-block.wgsl.jinja`](build/webgpu/datamove-split-block.wgsl.jinja)
60
  - [`split-n.wgsl.jinja`](build/webgpu/split-n.wgsl.jinja)
 
62
  ## Use with `@huggingface/kernels`
63
 
64
  ```sh
65
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
66
  ```
67
 
68
  Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
build/webgpu/datamove-split-block.wgsl.jinja CHANGED
@@ -48,11 +48,10 @@ fn copy_{{ name }}(group_index: u32) {
48
  if (lane_base + 2u < {{ run }}u) { {{ name }}[out_base + 2u] = input[in_base + 2u]; }
49
  if (lane_base + 3u < {{ run }}u) { {{ name }}[out_base + 3u] = input[in_base + 3u]; }
50
  }
51
- {%- endmacro %}
52
  {{ scalar_copy_fn("y0", run0_scalar, groups0, 0) }}
53
  {{ scalar_copy_fn("y1", run1_scalar, groups1, start1_scalar) }}
54
  {{ scalar_copy_fn("y2", run2_scalar, groups2, start2_scalar) }}
55
-
56
  @compute @workgroup_size({{ workgroupSize }})
57
  fn main(
58
  @builtin(global_invocation_id) gid: vec3<u32>,
 
48
  if (lane_base + 2u < {{ run }}u) { {{ name }}[out_base + 2u] = input[in_base + 2u]; }
49
  if (lane_base + 3u < {{ run }}u) { {{ name }}[out_base + 3u] = input[in_base + 3u]; }
50
  }
51
+ {% endmacro %}
52
  {{ scalar_copy_fn("y0", run0_scalar, groups0, 0) }}
53
  {{ scalar_copy_fn("y1", run1_scalar, groups1, start1_scalar) }}
54
  {{ scalar_copy_fn("y2", run2_scalar, groups2, start2_scalar) }}
 
55
  @compute @workgroup_size({{ workgroupSize }})
56
  fn main(
57
  @builtin(global_invocation_id) gid: vec3<u32>,
build/webgpu/manifest.json CHANGED
@@ -15,7 +15,7 @@
15
  },
16
  "attributes": { "axis": { "default": 0 }, "num_outputs": {} },
17
  "typeConstraints": {
18
- "T": ["float32", "float16", "uint32", "int32", "int16", "uint8", "int8", "bool"],
19
  "S": ["int64"]
20
  },
21
  "tunables": { "WORKGROUP_SIZE": { "default": 256 } },
@@ -34,32 +34,33 @@
34
  "twoBlockContract": "twoOutputContract and outer(shapes.y0, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y0, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y1, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y1, attrs.axis) == inner(shapes.input, attrs.axis)",
35
  "threeBlockContract": "threeOutputContract and outer(shapes.y0, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y0, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y1, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y1, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y2, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y2, attrs.axis) == inner(shapes.input, attrs.axis)",
36
  "fourBlockContract": "fourOutputContract and outer(shapes.y0, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y0, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y1, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y1, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y2, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y2, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y3, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y3, attrs.axis) == inner(shapes.input, attrs.axis)",
37
- "scalar": "dtypes.T"
 
 
 
38
  },
39
  "bindings": {
40
- "input": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
41
- "y0": { "buffer": "storage", "elementType": "$vectorScalar" },
42
- "y1": { "buffer": "storage", "elementType": "$vectorScalar" },
43
- "input_2": { "name": "input", "buffer": "read-only-storage", "elementType": "$ioElement" },
44
- "y0_2": { "name": "y0", "buffer": "storage", "elementType": "$ioElement" },
45
- "y1_2": { "name": "y1", "buffer": "storage", "elementType": "$ioElement" },
46
- "y2": { "buffer": "storage", "elementType": "$ioElement" },
47
- "input_3": { "name": "input", "buffer": "read-only-storage", "elementType": "$scalar" },
48
- "y0_3": { "name": "y0", "buffer": "storage", "elementType": "$scalar" },
49
- "y1_3": { "name": "y1", "buffer": "storage", "elementType": "$scalar" },
50
  "params": {
51
- "buffer": "uniform",
52
  "struct": [
53
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
54
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" }
55
  ]
56
  },
57
- "y2_3": { "name": "y2", "buffer": "storage", "elementType": "$scalar" },
58
- "y3_2": { "name": "y3", "buffer": "storage", "elementType": "$scalar" },
59
- "y4": { "buffer": "storage", "elementType": "$scalar" },
60
- "params_2": {
61
  "name": "params",
62
- "buffer": "uniform",
63
  "struct": [
64
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
65
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" },
@@ -68,9 +69,8 @@
68
  { "name": "y4Count", "type": "u32", "value": "numel(shapes.y4)" }
69
  ]
70
  },
71
- "params_3": {
72
  "name": "params",
73
- "buffer": "uniform",
74
  "struct": [
75
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
76
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" },
@@ -78,9 +78,8 @@
78
  { "name": "y3Count", "type": "u32", "value": "numel(shapes.y3)" }
79
  ]
80
  },
81
- "params_4": {
82
  "name": "params",
83
- "buffer": "uniform",
84
  "struct": [
85
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
86
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" },
@@ -96,7 +95,7 @@
96
  "passes": [
97
  {
98
  "id": "main",
99
- "name": "Split.copy",
100
  "shader": "datamove-flat-copy.wgsl.jinja",
101
  "derive": { "count": "numel(shapes.y0)" },
102
  "bindings": [
@@ -115,21 +114,14 @@
115
  {
116
  "id": "two_outputs_block_vec4",
117
  "priority": 15,
118
- "when": ["twoBlockContract", "(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis)) % 4 == 0", "(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis)) % 4 == 0"],
119
  "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
120
  "passes": [
121
  {
122
  "id": "main",
123
- "name": "Split.blockVec4",
124
  "shader": "datamove-split-block.wgsl.jinja",
125
- "derive": {
126
- "inputShape": "shapes.input",
127
- "y0Shape": "shapes.y0",
128
- "y1Shape": "shapes.y1",
129
- "rank": "ranks.input",
130
- "axisSpec": "axis",
131
- "outputCountSpec": 2
132
- },
133
  "bindings": ["input", "y0", "y1"],
134
  "dispatch": {
135
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1)) / 4), (workgroupSize)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
@@ -142,24 +134,21 @@
142
  {
143
  "id": "three_outputs_block_scalar_x4",
144
  "priority": 14,
145
- "when": ["threeBlockContract", "dim(shapes.y0, attrs.axis) > 0", "dim(shapes.y1, attrs.axis) > 0", "dim(shapes.y2, attrs.axis) > 0", "max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2)) >= 16"],
146
  "derive": { "ioElement": "dtypes.T" },
147
  "passes": [
148
  {
149
  "id": "main",
150
- "name": "Split3.blockScalarX4",
151
  "shader": "datamove-split-block.wgsl.jinja",
152
  "derive": {
153
- "inputShape": "shapes.input",
154
  "y0Shape": "shapes.y0",
155
  "y1Shape": "shapes.y1",
156
  "y2Shape": "shapes.y2",
157
- "rank": "ranks.input",
158
- "axisSpec": "axis",
159
  "outputCountSpec": 3,
160
  "scalarBoundX4": true
161
  },
162
- "bindings": ["input_2", "y0_2", "y1_2", "y2"],
163
  "dispatch": {
164
  "x": "min(ceilDiv((max(outer(shapes.y0, attrs.axis) * ceil(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis) / 4), outer(shapes.y1, attrs.axis) * ceil(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis) / 4), outer(shapes.y2, attrs.axis) * ceil(dim(shapes.y2, attrs.axis) * inner(shapes.y2, attrs.axis) / 4))), (workgroupSize)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
165
  "y": 1,
@@ -171,23 +160,15 @@
171
  {
172
  "id": "three_outputs_block_vec4",
173
  "priority": 15,
174
- "when": ["threeBlockContract", "(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis)) % 4 == 0", "(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis)) % 4 == 0", "(dim(shapes.y2, attrs.axis) * inner(shapes.y2, attrs.axis)) % 4 == 0"],
175
- "derive": { "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
176
  "passes": [
177
  {
178
  "id": "main",
179
- "name": "Split3.blockVec4",
180
  "shader": "datamove-split-block.wgsl.jinja",
181
- "derive": {
182
- "inputShape": "shapes.input",
183
- "y0Shape": "shapes.y0",
184
- "y1Shape": "shapes.y1",
185
- "y2Shape": "shapes.y2",
186
- "rank": "ranks.input",
187
- "axisSpec": "axis",
188
- "outputCountSpec": 3
189
- },
190
- "bindings": ["input_2", "y0_2", "y1_2", "y2"],
191
  "dispatch": {
192
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2)) / 4), (workgroupSize)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
193
  "y": 1,
@@ -199,21 +180,18 @@
199
  {
200
  "id": "four_outputs_block_vec4",
201
  "priority": 25,
202
- "when": ["fourBlockContract", "(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis)) % 4 == 0", "(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis)) % 4 == 0", "(dim(shapes.y2, attrs.axis) * inner(shapes.y2, attrs.axis)) % 4 == 0", "(dim(shapes.y3, attrs.axis) * inner(shapes.y3, attrs.axis)) % 4 == 0"],
203
  "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
204
  "passes": [
205
  {
206
  "id": "main",
207
- "name": "Split4.blockVec4",
208
  "shader": "datamove-split-block.wgsl.jinja",
209
  "derive": {
210
- "inputShape": "shapes.input",
211
  "y0Shape": "shapes.y0",
212
  "y1Shape": "shapes.y1",
213
  "y2Shape": "shapes.y2",
214
  "y3Shape": "shapes.y3",
215
- "rank": "ranks.input",
216
- "axisSpec": "axis",
217
  "outputCountSpec": 4
218
  },
219
  "bindings": [
@@ -240,12 +218,9 @@
240
  "name": "Split",
241
  "shader": "split-n.wgsl.jinja",
242
  "derive": {
243
- "inputShape": "shapes.input",
244
- "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1} ]",
245
- "rank": "ranks.input",
246
- "axisSpec": "axis"
247
  },
248
- "bindings": ["input_3", "y0_3", "y1_3", "params"],
249
  "dispatch": {
250
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1))), (workgroupSize)), 65535)",
251
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1))), (workgroupSize)), 65535)",
@@ -264,12 +239,9 @@
264
  "name": "SplitN",
265
  "shader": "split-n.wgsl.jinja",
266
  "derive": {
267
- "inputShape": "shapes.input",
268
- "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1}, {\"name\": \"y2\", \"shape\": shapes.y2}, {\"name\": \"y3\", \"shape\": shapes.y3}, {\"name\": \"y4\", \"shape\": shapes.y4} ]",
269
- "rank": "ranks.input",
270
- "axisSpec": "axis"
271
  },
272
- "bindings": ["input_3", "y0_3", "y1_3", "y2_3", "y3_2", "y4", "params_2"],
273
  "dispatch": {
274
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3), numel(shapes.y4))), (workgroupSize)), 65535)",
275
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3), numel(shapes.y4))), (workgroupSize)), 65535)",
@@ -288,12 +260,9 @@
288
  "name": "Split4",
289
  "shader": "split-n.wgsl.jinja",
290
  "derive": {
291
- "inputShape": "shapes.input",
292
- "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1}, {\"name\": \"y2\", \"shape\": shapes.y2}, {\"name\": \"y3\", \"shape\": shapes.y3} ]",
293
- "rank": "ranks.input",
294
- "axisSpec": "axis"
295
  },
296
- "bindings": ["input_3", "y0_3", "y1_3", "y2_3", "y3_2", "params_3"],
297
  "dispatch": {
298
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3))), (workgroupSize)), 65535)",
299
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3))), (workgroupSize)), 65535)",
@@ -312,12 +281,9 @@
312
  "name": "Split3",
313
  "shader": "split-n.wgsl.jinja",
314
  "derive": {
315
- "inputShape": "shapes.input",
316
- "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1}, {\"name\": \"y2\", \"shape\": shapes.y2} ]",
317
- "rank": "ranks.input",
318
- "axisSpec": "axis"
319
  },
320
- "bindings": ["input_3", "y0_3", "y1_3", "y2_3", "params_4"],
321
  "dispatch": {
322
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2))), (workgroupSize)), 65535)",
323
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2))), (workgroupSize)), 65535)",
 
15
  },
16
  "attributes": { "axis": { "default": 0 }, "num_outputs": {} },
17
  "typeConstraints": {
18
+ "T": ["float32", "float16", "uint32", "int32", "int16", "uint8", "int8", "bool", "int64"],
19
  "S": ["int64"]
20
  },
21
  "tunables": { "WORKGROUP_SIZE": { "default": 256 } },
 
34
  "twoBlockContract": "twoOutputContract and outer(shapes.y0, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y0, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y1, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y1, attrs.axis) == inner(shapes.input, attrs.axis)",
35
  "threeBlockContract": "threeOutputContract and outer(shapes.y0, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y0, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y1, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y1, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y2, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y2, attrs.axis) == inner(shapes.input, attrs.axis)",
36
  "fourBlockContract": "fourOutputContract and outer(shapes.y0, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y0, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y1, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y1, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y2, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y2, attrs.axis) == inner(shapes.input, attrs.axis) and outer(shapes.y3, attrs.axis) == outer(shapes.input, attrs.axis) and inner(shapes.y3, attrs.axis) == inner(shapes.input, attrs.axis)",
37
+ "scalar": "dtypes.T",
38
+ "inputShape": "shapes.input",
39
+ "rank": "ranks.input",
40
+ "axisSpec": "axis"
41
  },
42
  "bindings": {
43
+ "input": { "elementType": "$vectorScalar" },
44
+ "y0": { "elementType": "$vectorScalar" },
45
+ "y1": { "elementType": "$vectorScalar" },
46
+ "input_main": { "name": "input", "elementType": "$ioElement" },
47
+ "y0_main": { "name": "y0", "elementType": "$ioElement" },
48
+ "y1_main": { "name": "y1", "elementType": "$ioElement" },
49
+ "y2": { "elementType": "$ioElement" },
50
+ "input_scalar": { "name": "input", "elementType": "$scalar" },
51
+ "y0_scalar": { "name": "y0", "elementType": "$scalar" },
52
+ "y1_scalar": { "name": "y1", "elementType": "$scalar" },
53
  "params": {
 
54
  "struct": [
55
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
56
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" }
57
  ]
58
  },
59
+ "y2_main": { "name": "y2", "elementType": "$scalar" },
60
+ "y3_main": { "name": "y3", "elementType": "$scalar" },
61
+ "y4": { "elementType": "$scalar" },
62
+ "params_main": {
63
  "name": "params",
 
64
  "struct": [
65
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
66
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" },
 
69
  { "name": "y4Count", "type": "u32", "value": "numel(shapes.y4)" }
70
  ]
71
  },
72
+ "params_split_n": {
73
  "name": "params",
 
74
  "struct": [
75
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
76
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" },
 
78
  { "name": "y3Count", "type": "u32", "value": "numel(shapes.y3)" }
79
  ]
80
  },
81
+ "params__uniform": {
82
  "name": "params",
 
83
  "struct": [
84
  { "name": "y0Count", "type": "u32", "value": "numel(shapes.y0)" },
85
  { "name": "y1Count", "type": "u32", "value": "numel(shapes.y1)" },
 
95
  "passes": [
96
  {
97
  "id": "main",
98
+ "name": "Split.Copy",
99
  "shader": "datamove-flat-copy.wgsl.jinja",
100
  "derive": { "count": "numel(shapes.y0)" },
101
  "bindings": [
 
114
  {
115
  "id": "two_outputs_block_vec4",
116
  "priority": 15,
117
+ "when": ["dtypes.T != \"vec2<u32>\"", "twoBlockContract", "(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis)) % 4 == 0", "(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis)) % 4 == 0"],
118
  "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
119
  "passes": [
120
  {
121
  "id": "main",
122
+ "name": "Split.BlockVec4",
123
  "shader": "datamove-split-block.wgsl.jinja",
124
+ "derive": { "y0Shape": "shapes.y0", "y1Shape": "shapes.y1", "outputCountSpec": 2 },
 
 
 
 
 
 
 
125
  "bindings": ["input", "y0", "y1"],
126
  "dispatch": {
127
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1)) / 4), (workgroupSize)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
 
134
  {
135
  "id": "three_outputs_block_scalar_x4",
136
  "priority": 14,
137
+ "when": ["dtypes.T != \"vec2<u32>\"", "threeBlockContract", "dim(shapes.y0, attrs.axis) > 0", "dim(shapes.y1, attrs.axis) > 0", "dim(shapes.y2, attrs.axis) > 0", "max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2)) >= 16"],
138
  "derive": { "ioElement": "dtypes.T" },
139
  "passes": [
140
  {
141
  "id": "main",
142
+ "name": "Split3.BlockScalarX4",
143
  "shader": "datamove-split-block.wgsl.jinja",
144
  "derive": {
 
145
  "y0Shape": "shapes.y0",
146
  "y1Shape": "shapes.y1",
147
  "y2Shape": "shapes.y2",
 
 
148
  "outputCountSpec": 3,
149
  "scalarBoundX4": true
150
  },
151
+ "bindings": ["input_main", "y0_main", "y1_main", "y2"],
152
  "dispatch": {
153
  "x": "min(ceilDiv((max(outer(shapes.y0, attrs.axis) * ceil(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis) / 4), outer(shapes.y1, attrs.axis) * ceil(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis) / 4), outer(shapes.y2, attrs.axis) * ceil(dim(shapes.y2, attrs.axis) * inner(shapes.y2, attrs.axis) / 4))), (workgroupSize)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
154
  "y": 1,
 
160
  {
161
  "id": "three_outputs_block_vec4",
162
  "priority": 15,
163
+ "when": ["dtypes.T != \"vec2<u32>\"", "threeBlockContract", "(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis)) % 4 == 0", "(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis)) % 4 == 0", "(dim(shapes.y2, attrs.axis) * inner(shapes.y2, attrs.axis)) % 4 == 0"],
164
+ "derive": { "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" },
165
  "passes": [
166
  {
167
  "id": "main",
168
+ "name": "Split3.BlockVec4",
169
  "shader": "datamove-split-block.wgsl.jinja",
170
+ "derive": { "y0Shape": "shapes.y0", "y1Shape": "shapes.y1", "y2Shape": "shapes.y2", "outputCountSpec": 3 },
171
+ "bindings": ["input_main", "y0_main", "y1_main", "y2"],
 
 
 
 
 
 
 
 
172
  "dispatch": {
173
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2)) / 4), (workgroupSize)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
174
  "y": 1,
 
180
  {
181
  "id": "four_outputs_block_vec4",
182
  "priority": 25,
183
+ "when": ["dtypes.T != \"vec2<u32>\"", "fourBlockContract", "(dim(shapes.y0, attrs.axis) * inner(shapes.y0, attrs.axis)) % 4 == 0", "(dim(shapes.y1, attrs.axis) * inner(shapes.y1, attrs.axis)) % 4 == 0", "(dim(shapes.y2, attrs.axis) * inner(shapes.y2, attrs.axis)) % 4 == 0", "(dim(shapes.y3, attrs.axis) * inner(shapes.y3, attrs.axis)) % 4 == 0"],
184
  "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
185
  "passes": [
186
  {
187
  "id": "main",
188
+ "name": "Split4.BlockVec4",
189
  "shader": "datamove-split-block.wgsl.jinja",
190
  "derive": {
 
191
  "y0Shape": "shapes.y0",
192
  "y1Shape": "shapes.y1",
193
  "y2Shape": "shapes.y2",
194
  "y3Shape": "shapes.y3",
 
 
195
  "outputCountSpec": 4
196
  },
197
  "bindings": [
 
218
  "name": "Split",
219
  "shader": "split-n.wgsl.jinja",
220
  "derive": {
221
+ "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1} ]"
 
 
 
222
  },
223
+ "bindings": ["input_scalar", "y0_scalar", "y1_scalar", "params"],
224
  "dispatch": {
225
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1))), (workgroupSize)), 65535)",
226
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1))), (workgroupSize)), 65535)",
 
239
  "name": "SplitN",
240
  "shader": "split-n.wgsl.jinja",
241
  "derive": {
242
+ "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1}, {\"name\": \"y2\", \"shape\": shapes.y2}, {\"name\": \"y3\", \"shape\": shapes.y3}, {\"name\": \"y4\", \"shape\": shapes.y4} ]"
 
 
 
243
  },
244
+ "bindings": ["input_scalar", "y0_scalar", "y1_scalar", "y2_main", "y3_main", "y4", "params_main"],
245
  "dispatch": {
246
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3), numel(shapes.y4))), (workgroupSize)), 65535)",
247
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3), numel(shapes.y4))), (workgroupSize)), 65535)",
 
260
  "name": "Split4",
261
  "shader": "split-n.wgsl.jinja",
262
  "derive": {
263
+ "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1}, {\"name\": \"y2\", \"shape\": shapes.y2}, {\"name\": \"y3\", \"shape\": shapes.y3} ]"
 
 
 
264
  },
265
+ "bindings": ["input_scalar", "y0_scalar", "y1_scalar", "y2_main", "y3_main", "params_split_n"],
266
  "dispatch": {
267
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3))), (workgroupSize)), 65535)",
268
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2), numel(shapes.y3))), (workgroupSize)), 65535)",
 
281
  "name": "Split3",
282
  "shader": "split-n.wgsl.jinja",
283
  "derive": {
284
+ "outputs": "[ {\"name\": \"y0\", \"shape\": shapes.y0}, {\"name\": \"y1\", \"shape\": shapes.y1}, {\"name\": \"y2\", \"shape\": shapes.y2} ]"
 
 
 
285
  },
286
+ "bindings": ["input_scalar", "y0_scalar", "y1_scalar", "y2_main", "params__uniform"],
287
  "dispatch": {
288
  "x": "min(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2))), (workgroupSize)), 65535)",
289
  "y": "ceilDiv(ceilDiv((max(numel(shapes.y0), numel(shapes.y1), numel(shapes.y2))), (workgroupSize)), 65535)",
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.Split",
3
- "id": "_ai_onnx_split_webgpu_905cd69",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -9,15 +9,15 @@
9
  "files": {
10
  "bench.json": "k7jAXfnplRAPstbkQGbKdFPxoWX7R9DW+MK+4NiYTLU=",
11
  "datamove-flat-copy.wgsl.jinja": "g9d62mer2bmHfbScqX5CIT5Zq/Pdinac0W2dH0CD+5s=",
12
- "datamove-split-block.wgsl.jinja": "shWWQMUtwoSmMEgOn2yy1spChp86zzPkm/Cb7QiAnWs=",
13
- "manifest.json": "qnPxFTTkmu2fktjom8/oaTjlNXRcGD2kg+oeCsJ6RQc=",
14
- "split-n.wgsl.jinja": "z574+TJ8XGpxhV0wwxRfsHXSsSu7Ry4dvhXBxHcB6Cw=",
15
- "test.json": "71+a2DHyXRRHacOOABnDWrpe60/NRk7A/jP9UqPwpFE="
16
  }
17
  },
18
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
19
  "webgpu": {
20
- "manifestSpec": "2.0",
21
  "variants": {
22
  "one_output_copy": ["datamove-flat-copy.wgsl.jinja"],
23
  "two_outputs_block_vec4": ["datamove-split-block.wgsl.jinja"],
 
1
  {
2
  "name": "ai.onnx.Split",
3
+ "id": "_ai_onnx_split_webgpu_4487c2d",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
9
  "files": {
10
  "bench.json": "k7jAXfnplRAPstbkQGbKdFPxoWX7R9DW+MK+4NiYTLU=",
11
  "datamove-flat-copy.wgsl.jinja": "g9d62mer2bmHfbScqX5CIT5Zq/Pdinac0W2dH0CD+5s=",
12
+ "datamove-split-block.wgsl.jinja": "J9ZUsshWqXU6Y0INvyX/YN+WxVHLjbU+V7uRqVubCHo=",
13
+ "manifest.json": "OcyOVgNQihorEw/D6hDVzXXjTAkMxJnVqzJz7lwMrA4=",
14
+ "split-n.wgsl.jinja": "ong6jtDTsoitSqDf1/Z3iKG8wYLIBhmJOtCyPg4JUlQ=",
15
+ "test.json": "ekrJpXlDxovWHU6x8i5hhjvc68sXmA09kbfHJhk4JU0="
16
  }
17
  },
18
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
19
  "webgpu": {
20
+ "manifestSpec": "2.1",
21
  "variants": {
22
  "one_output_copy": ["datamove-flat-copy.wgsl.jinja"],
23
  "two_outputs_block_vec4": ["datamove-split-block.wgsl.jinja"],
build/webgpu/split-n.wgsl.jinja CHANGED
@@ -2,6 +2,11 @@
2
  // adding all preceding outputs' cumulative split-axis extent. One invocation
3
  // handles the same flat position across outputs, and each output writes only
4
  // when that position is within its own element count.
 
 
 
 
 
5
  {{ env.wgsl.resourceDeclarations }}
6
 
7
  {% for output in outputs %}
@@ -42,9 +47,7 @@ fn input_offset_{{ output.name }}({% if out_count.value > 0 %}out_index: u32{% e
42
  {% endfor %}
43
  @compute @workgroup_size({{ workgroupSize }})
44
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
45
- // 2D-folded flat index: gid.y carries the high bits past the
46
- // per-axis dispatch fold width.
47
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ workgroupSize }}u;
48
  {% for output in outputs %}
49
  {% set out_count = namespace(value=1) %}
50
  {% for d in output.shape %}
 
2
  // adding all preceding outputs' cumulative split-axis extent. One invocation
3
  // handles the same flat position across outputs, and each output writes only
4
  // when that position is within its own element count.
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 }};{% endmacro %}
10
  {{ env.wgsl.resourceDeclarations }}
11
 
12
  {% for output in outputs %}
 
47
  {% endfor %}
48
  @compute @workgroup_size({{ workgroupSize }})
49
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
50
+ {{ flat_index_2d(workgroupSize, "i", "") }}
 
 
51
  {% for output in outputs %}
52
  {% set out_count = namespace(value=1) %}
53
  {% for d in output.shape %}
build/webgpu/test.json CHANGED
@@ -1231,7 +1231,7 @@
1231
  {
1232
  "name": "three_outputs_negative_axis_unequal_mod4_tails",
1233
  "provenance": {
1234
- "notes": "Compact scalar three-output lock with multiple outer rows and segment lengths 6/5/3 (mod-4 tails 2/1/3). Negative axis normalization and unequal source offsets expose cross-segment or row-boundary corruption in the x4-coarsened path."
1235
  },
1236
  "attrs": { "axis": -1 },
1237
  "inputs": {
@@ -1265,7 +1265,7 @@
1265
  {
1266
  "name": "block_scalar_x4_three_outputs_y2_dominant_groups",
1267
  "provenance": {
1268
- "notes": "Splitting axis 1 into runs of 2, 3, and 6 scalars selects grouped scalar copying because no run is four-aligned. The third output needs two four-scalar groups per row and therefore determines the loop bound."
1269
  },
1270
  "attrs": { "axis": 1 },
1271
  "inputs": {
@@ -1280,7 +1280,7 @@
1280
  {
1281
  "name": "block_vec4_three_outputs_y2_dominant_count",
1282
  "provenance": {
1283
- "notes": "Three-output vec4 block split where the last chunk is the largest: runs 2/2/4 rows of inner 2 give 4/4/8 scalars, every boundary vec4-aligned, so Y2's vec4 run is twice Y0's and Y1's. That makes the third output's element count the strict maximum, the regime where the loop bound comes from Y2 rather than from the first two outputs. Plain ONNX Split semantics: axis=1 with split=[2,2,4] summing to the axis dim 8."
1284
  },
1285
  "attrs": { "axis": 1 },
1286
  "inputs": {
@@ -1331,6 +1331,102 @@
1331
  "y2": { "dtype": "uint32", "shape": [2, 4], "tolerance": 0 },
1332
  "y3": { "dtype": "uint32", "shape": [2, 0], "tolerance": 0 }
1333
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1334
  }
1335
  ]
1336
  }
 
1231
  {
1232
  "name": "three_outputs_negative_axis_unequal_mod4_tails",
1233
  "provenance": {
1234
+ "notes": "Split along axis -1 into segments 6, 5, and 3 (mod-4 remainders 2, 1, and 3) across 2x3 outer rows; unequal offsets and negative-axis normalization check that segment and row boundaries don't corrupt neighboring data."
1235
  },
1236
  "attrs": { "axis": -1 },
1237
  "inputs": {
 
1265
  {
1266
  "name": "block_scalar_x4_three_outputs_y2_dominant_groups",
1267
  "provenance": {
1268
+ "notes": "Splitting axis one into runs of two, three, and six scalars checks all three outputs and their partial four-element storage groups."
1269
  },
1270
  "attrs": { "axis": 1 },
1271
  "inputs": {
 
1280
  {
1281
  "name": "block_vec4_three_outputs_y2_dominant_count",
1282
  "provenance": {
1283
+ "notes": "Split axis 1 into sizes 2, 2, and 4 (summing to 8) makes the third output's element count (16) twice the first two (8 each); checks that per-output sizing follows the largest segment rather than the first ones listed."
1284
  },
1285
  "attrs": { "axis": 1 },
1286
  "inputs": {
 
1331
  "y2": { "dtype": "uint32", "shape": [2, 4], "tolerance": 0 },
1332
  "y3": { "dtype": "uint32", "shape": [2, 0], "tolerance": 0 }
1333
  }
1334
+ },
1335
+ {
1336
+ "name": "ort_split_axis0_vec4_via_trailing_dims",
1337
+ "provenance": {
1338
+ "source": "js/web/test/data/ops/split.jsonc",
1339
+ "test": "Split on Axis 0 - vec4 across trailing dimensions / T[2,4,3]"
1340
+ },
1341
+ "attrs": { "axis": 0 },
1342
+ "inputs": {
1343
+ "input": {
1344
+ "dtype": "float32",
1345
+ "shape": [2, 4, 3],
1346
+ "data": {
1347
+ "kind": "values",
1348
+ "values": { "$ref": "#/fixtureArrays/ort_axis2_equal_three_outputs_input_input" }
1349
+ }
1350
+ }
1351
+ },
1352
+ "outputs": {
1353
+ "y0": { "dtype": "float32", "shape": [1, 4, 3], "tolerance": 0 },
1354
+ "y1": { "dtype": "float32", "shape": [1, 4, 3], "tolerance": 0 }
1355
+ }
1356
+ },
1357
+ {
1358
+ "name": "ort_split_axis1_vec4_unequal_segments_multi_outer",
1359
+ "provenance": {
1360
+ "source": "js/web/test/data/ops/split.jsonc",
1361
+ "test": "Split on Axis 1 - vec4 with unequal segments across multiple outer slices / T[2,3,4]"
1362
+ },
1363
+ "attrs": { "axis": 1 },
1364
+ "inputs": {
1365
+ "input": {
1366
+ "dtype": "float32",
1367
+ "shape": [2, 3, 4],
1368
+ "data": {
1369
+ "kind": "values",
1370
+ "values": { "$ref": "#/fixtureArrays/ort_axis2_equal_three_outputs_input_input" }
1371
+ }
1372
+ }
1373
+ },
1374
+ "outputs": {
1375
+ "y0": { "dtype": "float32", "shape": [2, 1, 4], "tolerance": 0 },
1376
+ "y1": { "dtype": "float32", "shape": [2, 2, 4], "tolerance": 0 }
1377
+ }
1378
+ },
1379
+ {
1380
+ "name": "int64_full_range_max_arity_int32_positions",
1381
+ "provenance": {
1382
+ "notes": "Synthetic five-output int32 Split contract fixture; every bounded output position receives a distinct consecutive slice."
1383
+ },
1384
+ "attrs": { "axis": 0 },
1385
+ "inputs": {
1386
+ "input": {
1387
+ "dtype": "int64",
1388
+ "shape": [10],
1389
+ "data": {
1390
+ "kind": "cycle",
1391
+ "values": ["0", "4294967296", "8589934593", "-4294967297", "9223372036854775807", "-9223372036854775808", "9223372036854775806", "-1"]
1392
+ }
1393
+ }
1394
+ },
1395
+ "outputs": {
1396
+ "y0": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1397
+ "y1": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1398
+ "y2": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1399
+ "y3": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1400
+ "y4": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 }
1401
+ },
1402
+ "tolerance": 0,
1403
+ "relTolerance": 0
1404
+ },
1405
+ {
1406
+ "name": "int64_full_range_max_arity_int16_positions",
1407
+ "provenance": {
1408
+ "notes": "Synthetic five-output int16 Split contract fixture; every bounded output position receives a distinct consecutive slice."
1409
+ },
1410
+ "attrs": { "axis": 0 },
1411
+ "inputs": {
1412
+ "input": {
1413
+ "dtype": "int64",
1414
+ "shape": [10],
1415
+ "data": {
1416
+ "kind": "cycle",
1417
+ "values": ["4294967296", "8589934593", "-4294967297", "9223372036854775807", "-9223372036854775808", "9223372036854775806", "-1", "0"]
1418
+ }
1419
+ }
1420
+ },
1421
+ "outputs": {
1422
+ "y0": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1423
+ "y1": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1424
+ "y2": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1425
+ "y3": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 },
1426
+ "y4": { "dtype": "int64", "shape": [2], "tolerance": 0, "relTolerance": 0 }
1427
+ },
1428
+ "tolerance": 0,
1429
+ "relTolerance": 0
1430
  }
1431
  ]
1432
  }