sync 6fdf6301e2bb
Browse files- README.md +6 -6
- build/webgpu/bench.json +2 -2
- build/webgpu/manifest.json +8 -12
- build/webgpu/metadata.json +7 -7
- build/webgpu/test.json +105 -3
- build/webgpu/tile.wgsl.jinja +9 -6
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
|
| 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.
|
| 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: [
|
| 72 |
-
repeats: { data: repeatsData, shape: [
|
| 73 |
}, {
|
| 74 |
-
outputs: { output: { shape: [
|
| 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
|
| 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": "
|
| 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.
|
| 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": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
-
"bench.json": "
|
| 11 |
"datamove-tile-vec4.wgsl.jinja": "2+EFl6e/twKm/3oj2RBh+sx0EhYl/NNGGloW33ZprLs=",
|
| 12 |
-
"manifest.json": "
|
| 13 |
-
"test.json": "
|
| 14 |
-
"tile.wgsl.jinja": "
|
| 15 |
}
|
| 16 |
},
|
| 17 |
-
"provenance": { "kernel": { "sha": "
|
| 18 |
"webgpu": {
|
| 19 |
-
"manifestSpec": "2.
|
| 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)
|
| 418 |
-
"notes": "Spec-valid rank-7 Tile.
|
| 419 |
},
|
| 420 |
"inputs": {
|
| 421 |
"input": {
|
|
@@ -536,7 +536,7 @@
|
|
| 536 |
{
|
| 537 |
"name": "inner_vec4_lastaxis_repeat_mul4",
|
| 538 |
"provenance": {
|
| 539 |
-
"notes": "The innermost 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 |
-
|
| 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 |
}
|