sync 6fdf6301e2bb
Browse files- README.md +3 -3
- build/webgpu/manifest.json +6 -6
- build/webgpu/metadata.json +7 -7
- build/webgpu/test.json +1 -1
- build/webgpu/unary-scalar.wgsl.jinja +5 -9
- build/webgpu/unary-vec4.wgsl.jinja +19 -11
README.md
CHANGED
|
@@ -39,14 +39,14 @@ See the [ONNX `Sign` spec](https://onnx.ai/onnx/operators/onnx__Sign.html) for t
|
|
| 39 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 40 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 41 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 42 |
-
- [`bench.json`](build/webgpu/bench.json) — benchmark
|
| 43 |
- [`unary-scalar.wgsl.jinja`](build/webgpu/unary-scalar.wgsl.jinja)
|
| 44 |
- [`unary-vec4.wgsl.jinja`](build/webgpu/unary-vec4.wgsl.jinja)
|
| 45 |
|
| 46 |
## Use with `@huggingface/kernels`
|
| 47 |
|
| 48 |
```sh
|
| 49 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 50 |
```
|
| 51 |
|
| 52 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
@@ -60,5 +60,5 @@ Replace each `*Data` placeholder with a typed array containing the corresponding
|
|
| 60 |
import { getKernel } from "@huggingface/kernels";
|
| 61 |
|
| 62 |
const kernel = await getKernel("webgpu-kernels/ai.onnx.Sign", { version: 1 });
|
| 63 |
-
const { output } = await kernel({ input: { data: inputData, shape: [] } });
|
| 64 |
```
|
|
|
|
| 39 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 40 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 41 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 42 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 43 |
- [`unary-scalar.wgsl.jinja`](build/webgpu/unary-scalar.wgsl.jinja)
|
| 44 |
- [`unary-vec4.wgsl.jinja`](build/webgpu/unary-vec4.wgsl.jinja)
|
| 45 |
|
| 46 |
## Use with `@huggingface/kernels`
|
| 47 |
|
| 48 |
```sh
|
| 49 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 50 |
```
|
| 51 |
|
| 52 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 60 |
import { getKernel } from "@huggingface/kernels";
|
| 61 |
|
| 62 |
const kernel = await getKernel("webgpu-kernels/ai.onnx.Sign", { version: 1 });
|
| 63 |
+
const { output } = await kernel({ input: { data: inputData, shape: [3] } });
|
| 64 |
```
|
build/webgpu/manifest.json
CHANGED
|
@@ -6,20 +6,19 @@
|
|
| 6 |
"outputs": { "output": { "dtype": "T", "rank": "ranks.input", "shape": "shapes.input" } },
|
| 7 |
"typeConstraints": { "T": ["float32", "float16", "int32", "uint32", "int16", "int8", "uint8"] },
|
| 8 |
"tunables": { "WORKGROUP_SIZE": { "default": 256 } },
|
| 9 |
-
"derive": { "scalar": "dtypes.T"
|
| 10 |
"variants": [
|
| 11 |
{
|
| 12 |
"id": "same_layout_vec4",
|
| 13 |
"priority": 20,
|
| 14 |
"when": ["numel(shapes.input) > 0", "numel(shapes.input) % 4 == 0", "numel(shapes.input) == numel(shapes.output)", "f16Ok(dtypes.T)"],
|
| 15 |
-
"derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
|
| 16 |
"passes": [
|
| 17 |
{
|
| 18 |
"id": "main",
|
| 19 |
-
"name": "Sign.
|
| 20 |
"shader": "unary-vec4.wgsl.jinja",
|
| 21 |
"derive": {
|
| 22 |
-
"
|
| 23 |
"vec4PerThread": "4 if numel(shapes.output) * dtypeBytes(tensorDtypes.output) <= 16777216 else 1"
|
| 24 |
},
|
| 25 |
"bindings": [
|
|
@@ -43,7 +42,7 @@
|
|
| 43 |
"id": "main",
|
| 44 |
"name": "Sign",
|
| 45 |
"shader": "unary-scalar.wgsl.jinja",
|
| 46 |
-
"derive": { "
|
| 47 |
"bindings": [
|
| 48 |
"input",
|
| 49 |
"output",
|
|
@@ -57,5 +56,6 @@
|
|
| 57 |
}
|
| 58 |
]
|
| 59 |
}
|
| 60 |
-
]
|
|
|
|
| 61 |
}
|
|
|
|
| 6 |
"outputs": { "output": { "dtype": "T", "rank": "ranks.input", "shape": "shapes.input" } },
|
| 7 |
"typeConstraints": { "T": ["float32", "float16", "int32", "uint32", "int16", "int8", "uint8"] },
|
| 8 |
"tunables": { "WORKGROUP_SIZE": { "default": 256 } },
|
| 9 |
+
"derive": { "scalar": "dtypes.T" },
|
| 10 |
"variants": [
|
| 11 |
{
|
| 12 |
"id": "same_layout_vec4",
|
| 13 |
"priority": 20,
|
| 14 |
"when": ["numel(shapes.input) > 0", "numel(shapes.input) % 4 == 0", "numel(shapes.input) == numel(shapes.output)", "f16Ok(dtypes.T)"],
|
|
|
|
| 15 |
"passes": [
|
| 16 |
{
|
| 17 |
"id": "main",
|
| 18 |
+
"name": "Sign.Vec4",
|
| 19 |
"shader": "unary-vec4.wgsl.jinja",
|
| 20 |
"derive": {
|
| 21 |
+
"vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
|
| 22 |
"vec4PerThread": "4 if numel(shapes.output) * dtypeBytes(tensorDtypes.output) <= 16777216 else 1"
|
| 23 |
},
|
| 24 |
"bindings": [
|
|
|
|
| 42 |
"id": "main",
|
| 43 |
"name": "Sign",
|
| 44 |
"shader": "unary-scalar.wgsl.jinja",
|
| 45 |
+
"derive": { "itemsPerInvocation": 4 },
|
| 46 |
"bindings": [
|
| 47 |
"input",
|
| 48 |
"output",
|
|
|
|
| 56 |
}
|
| 57 |
]
|
| 58 |
}
|
| 59 |
+
],
|
| 60 |
+
"bindings": {}
|
| 61 |
}
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "ai.onnx.Sign",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
@@ -8,15 +8,15 @@
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "BeK/pm3J1N63lmR/lECkTp1EWNnUeMFlM0F3OdgdhPk=",
|
| 11 |
-
"manifest.json": "
|
| 12 |
-
"test.json": "
|
| 13 |
-
"unary-scalar.wgsl.jinja": "
|
| 14 |
-
"unary-vec4.wgsl.jinja": "
|
| 15 |
}
|
| 16 |
},
|
| 17 |
-
"provenance": { "kernel": { "sha": "
|
| 18 |
"webgpu": {
|
| 19 |
-
"manifestSpec": "2.
|
| 20 |
"variants": { "same_layout_vec4": ["unary-vec4.wgsl.jinja"], "elementwise": ["unary-scalar.wgsl.jinja"] }
|
| 21 |
}
|
| 22 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "ai.onnx.Sign",
|
| 3 |
+
"id": "_ai_onnx_sign_webgpu_40dc75e",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "BeK/pm3J1N63lmR/lECkTp1EWNnUeMFlM0F3OdgdhPk=",
|
| 11 |
+
"manifest.json": "s+ABXyddsttzhBW+Uw1Nhib/44ngfT0JQFJSN3mRH2I=",
|
| 12 |
+
"test.json": "eXJxP7h/I7/77rVPYQL1trwUfya7yb6ZZnUb0E3P0YI=",
|
| 13 |
+
"unary-scalar.wgsl.jinja": "YEs9JqiNGorr5F+CR1xAjYxHi5MEUt8vVhfOL/Mcy8E=",
|
| 14 |
+
"unary-vec4.wgsl.jinja": "l+T86ZnOFzPHh26bOLUMHmlsoptZzt6S87Qo1aQ6eao="
|
| 15 |
}
|
| 16 |
},
|
| 17 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 18 |
"webgpu": {
|
| 19 |
+
"manifestSpec": "2.1",
|
| 20 |
"variants": { "same_layout_vec4": ["unary-vec4.wgsl.jinja"], "elementwise": ["unary-scalar.wgsl.jinja"] }
|
| 21 |
}
|
| 22 |
}
|
build/webgpu/test.json
CHANGED
|
@@ -219,7 +219,7 @@
|
|
| 219 |
"provenance": {
|
| 220 |
"source": "onnxruntime/test/providers/cpu/math/sign_test.cc",
|
| 221 |
"test": "MathOpTest.Sign_float",
|
| 222 |
-
"notes": "Extends the ORT float Sign coverage with
|
| 223 |
},
|
| 224 |
"inputs": {
|
| 225 |
"input": {
|
|
|
|
| 219 |
"provenance": {
|
| 220 |
"source": "onnxruntime/test/providers/cpu/math/sign_test.cc",
|
| 221 |
"test": "MathOpTest.Sign_float",
|
| 222 |
+
"notes": "Extends the ORT float Sign coverage with NaN preservation verified against ONNX Runtime's CPU provider."
|
| 223 |
},
|
| 224 |
"inputs": {
|
| 225 |
"input": {
|
build/webgpu/unary-scalar.wgsl.jinja
CHANGED
|
@@ -12,15 +12,11 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
| 12 |
for (var i = begin; i < end; i = i + 1u) {
|
| 13 |
{%- endmacro %}
|
| 14 |
{% macro flat_tail_close() %}
|
| 15 |
-
}
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
// Scalar unary elementwise implementation. Specialization emits only the
|
| 19 |
-
// selected operation and any numerical helper it requires.
|
| 20 |
-
{% if usesF16 %}
|
| 21 |
-
enable f16;
|
| 22 |
-
{% endif %}
|
| 23 |
{{ env.wgsl.resourceDeclarations }}
|
|
|
|
| 24 |
{{ flat_tail_open() }}
|
| 25 |
{% if scalar == "i32" %}
|
| 26 |
let v = input[i];
|
|
@@ -38,5 +34,5 @@ enable f16;
|
|
| 38 |
if (v_is_nan) { s = v; }
|
| 39 |
output[i] = {{ scalar }}(s);
|
| 40 |
{% endif %}
|
| 41 |
-
{{ flat_tail_close()
|
| 42 |
}
|
|
|
|
| 12 |
for (var i = begin; i < end; i = i + 1u) {
|
| 13 |
{%- endmacro %}
|
| 14 |
{% macro flat_tail_close() %}
|
| 15 |
+
}{% endmacro %}
|
| 16 |
+
// Scalar unary elementwise kernel: applies the selected operation, with the
|
| 17 |
+
// numerical helpers it uses, to each element.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
{{ env.wgsl.resourceDeclarations }}
|
| 19 |
+
|
| 20 |
{{ flat_tail_open() }}
|
| 21 |
{% if scalar == "i32" %}
|
| 22 |
let v = input[i];
|
|
|
|
| 34 |
if (v_is_nan) { s = v; }
|
| 35 |
output[i] = {{ scalar }}(s);
|
| 36 |
{% endif %}
|
| 37 |
+
{{ flat_tail_close() }}
|
| 38 |
}
|
build/webgpu/unary-vec4.wgsl.jinja
CHANGED
|
@@ -2,34 +2,42 @@
|
|
| 2 |
// component.
|
| 3 |
{{ env.wgsl.resourceDeclarations }}
|
| 4 |
|
| 5 |
-
|
| 6 |
-
{% set
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
{% if vec4PerThread > 1 %}
|
| 8 |
const ITEMS: u32 = {{ vec4PerThread }}u;
|
| 9 |
-
{% endif %}
|
| 10 |
|
|
|
|
| 11 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 12 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 13 |
-
// 2D-folded flat index: gid.y carries the high bits when the element count exceeds the
|
| 14 |
-
// per-axis dispatch fold width (the dispatch caps x and spills the rest into y).
|
| 15 |
{% if vec4PerThread > 1 %}
|
| 16 |
// Each invocation walks ITEMS vec4 groups a span apart. Consecutive lanes
|
| 17 |
// access consecutive words on every step, while each lane can keep several
|
| 18 |
// independent loads in flight.
|
| 19 |
-
|
| 20 |
let span = (params.count + ITEMS - 1u) / ITEMS;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
for (var j = 0u; j < ITEMS; j = j + 1u) {
|
| 22 |
let i = tid + j * span;
|
| 23 |
if (i >= params.count) {
|
| 24 |
break;
|
| 25 |
}
|
| 26 |
{% else %}
|
| 27 |
-
|
| 28 |
-
if (i >= params.count) {
|
| 29 |
-
return;
|
| 30 |
-
}
|
| 31 |
{% endif %}
|
| 32 |
-
|
| 33 |
let xv = x[i];
|
| 34 |
{% if scalar == "i32" %}
|
| 35 |
y[i] = select(select(vec4<i32>(0i), vec4<i32>(1i), xv > vec4<i32>(0i)), vec4<i32>(-1i), xv < vec4<i32>(0i));
|
|
|
|
| 2 |
// component.
|
| 3 |
{{ env.wgsl.resourceDeclarations }}
|
| 4 |
|
| 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 |
+
{% if bound == "" %}
|
| 10 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% else %}
|
| 11 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
|
| 12 |
+
if ({{ name }} >= {{ bound }}) {
|
| 13 |
+
return;
|
| 14 |
+
}{% endif %}{% endmacro %}
|
| 15 |
{% if vec4PerThread > 1 %}
|
| 16 |
const ITEMS: u32 = {{ vec4PerThread }}u;
|
|
|
|
| 17 |
|
| 18 |
+
{% endif %}
|
| 19 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 20 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
|
|
|
|
|
| 21 |
{% if vec4PerThread > 1 %}
|
| 22 |
// Each invocation walks ITEMS vec4 groups a span apart. Consecutive lanes
|
| 23 |
// access consecutive words on every step, while each lane can keep several
|
| 24 |
// independent loads in flight.
|
| 25 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE, "tid", "") }}
|
| 26 |
let span = (params.count + ITEMS - 1u) / ITEMS;
|
| 27 |
+
// Lanes at or past span exist only because the dispatch rounds up to whole
|
| 28 |
+
// workgroups. Lane span + k would start on lane k's second group and rewrite
|
| 29 |
+
// up to ITEMS - 1 groups another lane already stored.
|
| 30 |
+
if (tid >= span) {
|
| 31 |
+
return;
|
| 32 |
+
}
|
| 33 |
for (var j = 0u; j < ITEMS; j = j + 1u) {
|
| 34 |
let i = tid + j * span;
|
| 35 |
if (i >= params.count) {
|
| 36 |
break;
|
| 37 |
}
|
| 38 |
{% else %}
|
| 39 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE) }}
|
|
|
|
|
|
|
|
|
|
| 40 |
{% endif %}
|
|
|
|
| 41 |
let xv = x[i];
|
| 42 |
{% if scalar == "i32" %}
|
| 43 |
y[i] = select(select(vec4<i32>(0i), vec4<i32>(1i), xv > vec4<i32>(0i)), vec4<i32>(-1i), xv < vec4<i32>(0i));
|