Xenova HF Staff commited on
Commit
76157fd
·
verified ·
1 Parent(s): 271fadd

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -33,7 +33,7 @@ See the [ONNX `Tile` spec](https://onnx.ai/onnx/operators/onnx__Tile.html) for t
33
 
34
  | Variable | Allowed dtypes |
35
  | --- | --- |
36
- | `T` | `float32`, `float16`, `uint32`, `int32`, `int16`, `uint8`, `int8`, `bool` |
37
  | `S` | `int64` |
38
 
39
  ## Files
@@ -41,14 +41,14 @@ See the [ONNX `Tile` spec](https://onnx.ai/onnx/operators/onnx__Tile.html) for t
41
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
42
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
43
  - [`test.json`](build/webgpu/test.json) — correctness cases
44
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
45
  - [`datamove-tile-vec4.wgsl.jinja`](build/webgpu/datamove-tile-vec4.wgsl.jinja)
46
  - [`tile.wgsl.jinja`](build/webgpu/tile.wgsl.jinja)
47
 
48
  ## Use with `@huggingface/kernels`
49
 
50
  ```sh
51
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
52
  ```
53
 
54
  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.
@@ -68,9 +68,9 @@ import { getKernel } from "@huggingface/kernels";
68
  const kernel = await getKernel("webgpu-kernels/ai.onnx.Tile", { version: 1 });
69
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
70
  const { output } = await kernel({
71
- input: { data: inputData, shape: [2, 1, 3] },
72
- repeats: { data: repeatsData, shape: [3] },
73
  }, {
74
- outputs: { output: { shape: [2, 1, 3], dtype: "float32" } },
75
  });
76
  ```
 
33
 
34
  | Variable | Allowed dtypes |
35
  | --- | --- |
36
+ | `T` | `float32`, `float16`, `uint32`, `int32`, `int16`, `uint8`, `int8`, `bool`, `int64` |
37
  | `S` | `int64` |
38
 
39
  ## Files
 
41
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
42
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
43
  - [`test.json`](build/webgpu/test.json) — correctness cases
44
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
45
  - [`datamove-tile-vec4.wgsl.jinja`](build/webgpu/datamove-tile-vec4.wgsl.jinja)
46
  - [`tile.wgsl.jinja`](build/webgpu/tile.wgsl.jinja)
47
 
48
  ## Use with `@huggingface/kernels`
49
 
50
  ```sh
51
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
52
  ```
53
 
54
  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.
 
68
  const kernel = await getKernel("webgpu-kernels/ai.onnx.Tile", { version: 1 });
69
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
70
  const { output } = await kernel({
71
+ input: { data: inputData, shape: [3] },
72
+ repeats: { data: repeatsData, shape: [1] },
73
  }, {
74
+ outputs: { output: { shape: [6], dtype: "int64" } },
75
  });
76
  ```
build/webgpu/bench.json CHANGED
@@ -93,7 +93,7 @@
93
  "name": "tile-f32-8x1x197x197-repeat-heads-generic-pathology",
94
  "preset": "stress",
95
  "provenance": {
96
- "notes": "A [B,1,S,S] attention mask tiled over 8 heads at the ViT-B/16 token count: the trailing unrepeated block is 197*197 = 38809 elements, not a multiple of four, so tileVec4Ok fails and the 2.5M-element copy runs the generic per-element route. Control: tile-f32-8x1x200x200-repeat-heads-control."
97
  },
98
  "vars": { "dtype": "float32", "outCount": 2483776 },
99
  "inputs": {
@@ -107,7 +107,7 @@
107
  "name": "tile-f32-8x1x200x200-repeat-heads-control",
108
  "preset": "stress",
109
  "provenance": {
110
- "notes": "Nearest multiple-of-4 token count for tile-f32-8x1x197x197-repeat-heads-generic-pathology (3% more bytes): the trailing block 200*200 folds to a vec4 tail and inner_vec4 is selected."
111
  },
112
  "vars": { "dtype": "float32", "outCount": 2560000 },
113
  "inputs": {
 
93
  "name": "tile-f32-8x1x197x197-repeat-heads-generic-pathology",
94
  "preset": "stress",
95
  "provenance": {
96
+ "notes": "A [B,1,S,S] attention mask repeated over eight heads at the ViT-B/16 token count: the trailing block has 197*197 elements, not divisible by four. The paired 200-token case checks an aligned block."
97
  },
98
  "vars": { "dtype": "float32", "outCount": 2483776 },
99
  "inputs": {
 
107
  "name": "tile-f32-8x1x200x200-repeat-heads-control",
108
  "preset": "stress",
109
  "provenance": {
110
+ "notes": "The next multiple-of-four token count above the paired 197-token case, with about 3% more bytes."
111
  },
112
  "vars": { "dtype": "float32", "outCount": 2560000 },
113
  "inputs": {
build/webgpu/manifest.json CHANGED
@@ -8,7 +8,7 @@
8
  },
9
  "outputs": { "output": { "dtype": "T", "rank": "ranks.input" } },
10
  "typeConstraints": {
11
- "T": ["float32", "float16", "uint32", "int32", "int16", "uint8", "int8", "bool"],
12
  "S": ["int64"]
13
  },
14
  "tunables": { "WORKGROUP_SIZE": { "default": 256 } },
@@ -16,27 +16,24 @@
16
  "tileTailAxes": "pick([[numel(suffix(shapes.input, ranks.input - 8)) == numel(suffix(shapes.output, ranks.input - 8)), 8], [numel(suffix(shapes.input, ranks.input - 7)) == numel(suffix(shapes.output, ranks.input - 7)), 7], [numel(suffix(shapes.input, ranks.input - 6)) == numel(suffix(shapes.output, ranks.input - 6)), 6], [numel(suffix(shapes.input, ranks.input - 5)) == numel(suffix(shapes.output, ranks.input - 5)), 5], [numel(suffix(shapes.input, ranks.input - 4)) == numel(suffix(shapes.output, ranks.input - 4)), 4], [numel(suffix(shapes.input, ranks.input - 3)) == numel(suffix(shapes.output, ranks.input - 3)), 3], [numel(suffix(shapes.input, ranks.input - 2)) == numel(suffix(shapes.output, ranks.input - 2)), 2], [numel(suffix(shapes.input, ranks.input - 1)) == numel(suffix(shapes.output, ranks.input - 1)), 1]], 0)",
17
  "tileFoldAxes": "min(tileTailAxes, ranks.input)",
18
  "tileTail": "numel(suffix(shapes.input, ranks.input - tileFoldAxes)) if tileFoldAxes > 0 else (dim(shapes.input, ranks.input - 1) if ranks.input > 0 else 0)",
19
- "tileVec4Ok": "ranks.input >= 1 and numel(shapes.output) > 0 and tileTail % 4 == 0"
 
 
 
20
  },
21
  "when": ["ranks.input == ranks.output", "ranks.repeats == 1", "dim(shapes.repeats, 0) == ranks.input", "f16Ok(dtypes.T)"],
22
  "variants": [
23
  {
24
  "id": "inner_vec4",
25
  "priority": 15,
26
- "when": ["ranks.input >= 1", "numel(shapes.output) > 0", "tileVec4Ok"],
27
  "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
28
  "passes": [
29
  {
30
  "id": "main",
31
- "name": "Tile.innerVec4",
32
  "shader": "datamove-tile-vec4.wgsl.jinja",
33
- "derive": {
34
- "inputShape": "shapes.input",
35
- "outputShape": "shapes.output",
36
- "rank": "ranks.input",
37
- "count": "numel(shapes.output) / 4",
38
- "foldAxes": "tileFoldAxes"
39
- },
40
  "bindings": [
41
  { "arg": "input", "elementType": "$vectorScalar" },
42
  { "arg": "output", "elementType": "$vectorScalar" }
@@ -57,7 +54,6 @@
57
  "id": "main",
58
  "name": "Tile",
59
  "shader": "tile.wgsl.jinja",
60
- "derive": { "inputShape": "shapes.input", "outputShape": "shapes.output", "rank": "ranks.input" },
61
  "bindings": [
62
  { "arg": "input", "elementType": "$scalar" },
63
  { "arg": "output", "elementType": "$scalar" },
 
8
  },
9
  "outputs": { "output": { "dtype": "T", "rank": "ranks.input" } },
10
  "typeConstraints": {
11
+ "T": ["float32", "float16", "uint32", "int32", "int16", "uint8", "int8", "bool", "int64"],
12
  "S": ["int64"]
13
  },
14
  "tunables": { "WORKGROUP_SIZE": { "default": 256 } },
 
16
  "tileTailAxes": "pick([[numel(suffix(shapes.input, ranks.input - 8)) == numel(suffix(shapes.output, ranks.input - 8)), 8], [numel(suffix(shapes.input, ranks.input - 7)) == numel(suffix(shapes.output, ranks.input - 7)), 7], [numel(suffix(shapes.input, ranks.input - 6)) == numel(suffix(shapes.output, ranks.input - 6)), 6], [numel(suffix(shapes.input, ranks.input - 5)) == numel(suffix(shapes.output, ranks.input - 5)), 5], [numel(suffix(shapes.input, ranks.input - 4)) == numel(suffix(shapes.output, ranks.input - 4)), 4], [numel(suffix(shapes.input, ranks.input - 3)) == numel(suffix(shapes.output, ranks.input - 3)), 3], [numel(suffix(shapes.input, ranks.input - 2)) == numel(suffix(shapes.output, ranks.input - 2)), 2], [numel(suffix(shapes.input, ranks.input - 1)) == numel(suffix(shapes.output, ranks.input - 1)), 1]], 0)",
17
  "tileFoldAxes": "min(tileTailAxes, ranks.input)",
18
  "tileTail": "numel(suffix(shapes.input, ranks.input - tileFoldAxes)) if tileFoldAxes > 0 else (dim(shapes.input, ranks.input - 1) if ranks.input > 0 else 0)",
19
+ "tileVec4Ok": "ranks.input >= 1 and numel(shapes.output) > 0 and tileTail % 4 == 0",
20
+ "inputShape": "shapes.input",
21
+ "outputShape": "shapes.output",
22
+ "rank": "ranks.input"
23
  },
24
  "when": ["ranks.input == ranks.output", "ranks.repeats == 1", "dim(shapes.repeats, 0) == ranks.input", "f16Ok(dtypes.T)"],
25
  "variants": [
26
  {
27
  "id": "inner_vec4",
28
  "priority": 15,
29
+ "when": ["dtypes.T != \"vec2<u32>\"", "ranks.input >= 1", "numel(shapes.output) > 0", "tileVec4Ok"],
30
  "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
31
  "passes": [
32
  {
33
  "id": "main",
34
+ "name": "Tile.InnerVec4",
35
  "shader": "datamove-tile-vec4.wgsl.jinja",
36
+ "derive": { "count": "numel(shapes.output) / 4", "foldAxes": "tileFoldAxes" },
 
 
 
 
 
 
37
  "bindings": [
38
  { "arg": "input", "elementType": "$vectorScalar" },
39
  { "arg": "output", "elementType": "$vectorScalar" }
 
54
  "id": "main",
55
  "name": "Tile",
56
  "shader": "tile.wgsl.jinja",
 
57
  "bindings": [
58
  { "arg": "input", "elementType": "$scalar" },
59
  { "arg": "output", "elementType": "$scalar" },
build/webgpu/metadata.json CHANGED
@@ -1,22 +1,22 @@
1
  {
2
  "name": "ai.onnx.Tile",
3
- "id": "_ai_onnx_tile_webgpu_f5f5692",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "YfCc57C/DdUbEVmhWy5t2xoP2Le9Yi9K1tINlEYgeLw=",
11
  "datamove-tile-vec4.wgsl.jinja": "2+EFl6e/twKm/3oj2RBh+sx0EhYl/NNGGloW33ZprLs=",
12
- "manifest.json": "YC55CwcGcRBjXpT8bLDiPA2dt7Ml7TBb2Z1w4Bmilos=",
13
- "test.json": "b8XhMm/6K0JPfgbWB718034tNom6lm9mc2dOE/ctr3U=",
14
- "tile.wgsl.jinja": "JTi9v0O2HGR8u6QvjZFSHW2cycDILumledR4r/GmkA8="
15
  }
16
  },
17
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
18
  "webgpu": {
19
- "manifestSpec": "2.0",
20
  "variants": { "inner_vec4": ["datamove-tile-vec4.wgsl.jinja"], "generic": ["tile.wgsl.jinja"] }
21
  }
22
  }
 
1
  {
2
  "name": "ai.onnx.Tile",
3
+ "id": "_ai_onnx_tile_webgpu_32e0f56",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "1eM36L7PfUEX11p0lx+TyZ25BWbeAJD+H3gHG3sSRZM=",
11
  "datamove-tile-vec4.wgsl.jinja": "2+EFl6e/twKm/3oj2RBh+sx0EhYl/NNGGloW33ZprLs=",
12
+ "manifest.json": "1KJxVLG3UNqIFOiQACZOvKRxTS8R046eRe9Hw0jYJUQ=",
13
+ "test.json": "Ur/BnRjKBHpjTa1kPqGIHWhfmOX42bRXne1iRh3xGas=",
14
+ "tile.wgsl.jinja": "/Se423/ugSXxY/IqEB7Ktbc7G3xf9sVsE13TSM419hE="
15
  }
16
  },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
  "webgpu": {
19
+ "manifestSpec": "2.1",
20
  "variants": { "inner_vec4": ["datamove-tile-vec4.wgsl.jinja"], "generic": ["tile.wgsl.jinja"] }
21
  }
22
  }
build/webgpu/test.json CHANGED
@@ -414,8 +414,8 @@
414
  {
415
  "name": "rank7_unit_input_alternating_repeats",
416
  "provenance": {
417
- "source": "ONNX Tile spec (no rank limit) + onnxruntime CPU Tile kernel",
418
- "notes": "Spec-valid rank-7 Tile. ORT CPU computes it via its generic strided copy."
419
  },
420
  "inputs": {
421
  "input": {
@@ -536,7 +536,7 @@
536
  {
537
  "name": "inner_vec4_lastaxis_repeat_mul4",
538
  "provenance": {
539
- "notes": "The innermost axis itself is repeated; its input extent (8) is a multiple of four so every output vec4 still lies inside one copy and the vec4 route applies the modulo on the last axis."
540
  },
541
  "inputs": {
542
  "input": {
@@ -611,6 +611,108 @@
611
  "repeats": { "dtype": "uint32", "shape": [3], "data": { "kind": "values", "values": [2, 1, 1] } }
612
  },
613
  "outputs": { "output": { "dtype": "float32", "shape": [4, 3, 3], "tolerance": 0 } }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
614
  }
615
  ]
616
  }
 
414
  {
415
  "name": "rank7_unit_input_alternating_repeats",
416
  "provenance": {
417
+ "source": "ONNX Tile spec (no rank limit) and the Tile kernel in ONNX Runtime's CPU provider",
418
+ "notes": "Spec-valid rank-7 Tile. ONNX Runtime's CPU provider computes it with its generic strided copy."
419
  },
420
  "inputs": {
421
  "input": {
 
536
  {
537
  "name": "inner_vec4_lastaxis_repeat_mul4",
538
  "provenance": {
539
+ "notes": "The innermost axis has extent eight and repeats in four-aligned copies, checking modulo wrap at each copy boundary."
540
  },
541
  "inputs": {
542
  "input": {
 
611
  "repeats": { "dtype": "uint32", "shape": [3], "data": { "kind": "values", "values": [2, 1, 1] } }
612
  },
613
  "outputs": { "output": { "dtype": "float32", "shape": [4, 3, 3], "tolerance": 0 } }
614
+ },
615
+ {
616
+ "name": "int64_full_range_rank2",
617
+ "inputs": {
618
+ "input": {
619
+ "dtype": "int64",
620
+ "shape": [2, 2],
621
+ "data": {
622
+ "kind": "cycle",
623
+ "values": ["0", "4294967296", "8589934593", "-4294967297", "9223372036854775807", "-9223372036854775808", "9223372036854775806", "-1"]
624
+ }
625
+ },
626
+ "repeats": { "dtype": "uint32", "shape": [2], "data": { "kind": "values", "values": [2, 3] } }
627
+ },
628
+ "outputs": { "output": { "dtype": "int64", "shape": [4, 6], "tolerance": 0, "relTolerance": 0 } },
629
+ "tolerance": 0,
630
+ "relTolerance": 0
631
+ },
632
+ {
633
+ "name": "int64_full_range_ort_float_2d_two_axes_full_repeat",
634
+ "provenance": {
635
+ "source": "onnxruntime/test/providers/cpu/tensor/tile_op_test.cc",
636
+ "test": "TensorOpTest.TileFloatType",
637
+ "notes": "RunTest<float>({2, 2}, {2, 2}) from ORT's shared Tile wrapper."
638
+ },
639
+ "inputs": {
640
+ "input": {
641
+ "dtype": "int64",
642
+ "shape": [2, 2],
643
+ "data": {
644
+ "kind": "cycle",
645
+ "values": ["4294967296", "8589934593", "-4294967297", "9223372036854775807", "-9223372036854775808", "9223372036854775806", "-1", "0"]
646
+ }
647
+ },
648
+ "repeats": { "dtype": "uint32", "shape": [2], "data": { "kind": "values", "values": [2, 2] } }
649
+ },
650
+ "outputs": { "output": { "dtype": "int64", "shape": [4, 4], "tolerance": 0, "relTolerance": 0 } },
651
+ "tolerance": 0,
652
+ "relTolerance": 0
653
+ },
654
+ {
655
+ "name": "ort_int64_1d",
656
+ "inputs": {
657
+ "input": {
658
+ "dtype": "int64",
659
+ "shape": [3],
660
+ "data": { "kind": "values", "values": ["281483566841860", "1407400653815816", "-1"] }
661
+ },
662
+ "repeats": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [2] } }
663
+ },
664
+ "outputs": {
665
+ "output": {
666
+ "dtype": "int64",
667
+ "shape": [6],
668
+ "data": {
669
+ "kind": "values",
670
+ "values": ["281483566841860", "1407400653815816", "-1", "281483566841860", "1407400653815816", "-1"]
671
+ }
672
+ }
673
+ }
674
+ },
675
+ {
676
+ "name": "ort_int64_2d",
677
+ "inputs": {
678
+ "input": {
679
+ "dtype": "int64",
680
+ "shape": [2, 2],
681
+ "data": { "kind": "values", "values": ["281483566841860", "1407400653815816", "-1", "281483566841860"] }
682
+ },
683
+ "repeats": { "dtype": "uint32", "shape": [2], "data": { "kind": "values", "values": [2, 1] } }
684
+ },
685
+ "outputs": {
686
+ "output": {
687
+ "dtype": "int64",
688
+ "shape": [4, 2],
689
+ "data": {
690
+ "kind": "values",
691
+ "values": ["281483566841860", "1407400653815816", "-1", "281483566841860", "281483566841860", "1407400653815816", "-1", "281483566841860"]
692
+ }
693
+ }
694
+ }
695
+ },
696
+ {
697
+ "name": "ort_int64_3d",
698
+ "inputs": {
699
+ "input": {
700
+ "dtype": "int64",
701
+ "shape": [1, 1, 3],
702
+ "data": { "kind": "values", "values": ["281483566841860", "1407400653815816", "-1"] }
703
+ },
704
+ "repeats": { "dtype": "uint32", "shape": [3], "data": { "kind": "values", "values": [1, 2, 1] } }
705
+ },
706
+ "outputs": {
707
+ "output": {
708
+ "dtype": "int64",
709
+ "shape": [1, 2, 3],
710
+ "data": {
711
+ "kind": "values",
712
+ "values": ["281483566841860", "1407400653815816", "-1", "281483566841860", "1407400653815816", "-1"]
713
+ }
714
+ }
715
+ }
716
  }
717
  ]
718
  }
build/webgpu/tile.wgsl.jinja CHANGED
@@ -1,3 +1,11 @@
 
 
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  fn input_offset({% if rank > 0 %}out_index: u32{% endif %}) -> u32 {
@@ -25,11 +33,6 @@ fn input_offset({% if rank > 0 %}out_index: u32{% endif %}) -> u32 {
25
 
26
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
27
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
28
- // 2D-folded flat index: gid.y carries the high bits past the
29
- // per-axis dispatch fold width.
30
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
31
- if (i >= params.count) {
32
- return;
33
- }
34
  output[i] = input[input_offset({% if rank > 0 %}i{% endif %})];
35
  }
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
6
+ if ({{ name }} >= {{ bound }}) {
7
+ return;
8
+ }{% endmacro %}
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
  fn input_offset({% if rank > 0 %}out_index: u32{% endif %}) -> u32 {
 
33
 
34
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
35
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
36
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE) }}
 
 
 
 
 
37
  output[i] = input[input_offset({% if rank > 0 %}i{% endif %})];
38
  }