sync 6fdf6301e2bb
Browse files- README.md +2 -2
- build/webgpu/bias-add.wgsl.jinja +7 -34
- build/webgpu/manifest.json +3 -3
- build/webgpu/metadata.json +5 -5
README.md
CHANGED
|
@@ -41,13 +41,13 @@ See the [ONNX Runtime `BiasAdd` contrib-operator spec](https://github.com/micros
|
|
| 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 |
- [`bias-add.wgsl.jinja`](build/webgpu/bias-add.wgsl.jinja)
|
| 46 |
|
| 47 |
## Use with `@huggingface/kernels`
|
| 48 |
|
| 49 |
```sh
|
| 50 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 51 |
```
|
| 52 |
|
| 53 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 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 |
- [`bias-add.wgsl.jinja`](build/webgpu/bias-add.wgsl.jinja)
|
| 46 |
|
| 47 |
## Use with `@huggingface/kernels`
|
| 48 |
|
| 49 |
```sh
|
| 50 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 51 |
```
|
| 52 |
|
| 53 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
build/webgpu/bias-add.wgsl.jinja
CHANGED
|
@@ -12,42 +12,15 @@ 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 |
-
{% macro flat_index_2d(name="i", bound="params.count", guardInline=false, note="dispatch-limit") %}
|
| 19 |
-
{% if note == "dispatch-limit" %}
|
| 20 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 21 |
-
// per-axis workgroup fold width (outputs > 16.7M elements).
|
| 22 |
-
{% elif note == "limit" %}
|
| 23 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 24 |
-
// per-axis workgroup fold width.
|
| 25 |
-
{% elif note == "device-axis" %}
|
| 26 |
-
// The flat dispatch is folded across x/y at a fixed per-axis workgroup
|
| 27 |
-
// width; gid.y carries the high portion of the output index.
|
| 28 |
-
{% elif note == "vec4-limit" %}
|
| 29 |
-
// 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
|
| 30 |
-
// per-axis workgroup fold width (the dispatch caps x and spills into y).
|
| 31 |
-
{% elif note == "element-limit" %}
|
| 32 |
-
// 2D-folded flat element index: gid.y carries the high bits past the
|
| 33 |
-
// dispatch's per-axis workgroup fold width.
|
| 34 |
-
{% elif note == "dispatch" %}
|
| 35 |
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 36 |
// per-axis workgroup fold width.
|
| 37 |
-
{
|
| 38 |
-
{% if bound == "" %}
|
| 39 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 40 |
-
{%- elif guardInline %}
|
| 41 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 42 |
-
if ({{ name }} >= {{ bound }}) { return; }
|
| 43 |
-
{%- else %}
|
| 44 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 45 |
if ({{ name }} >= {{ bound }}) {
|
| 46 |
return;
|
| 47 |
-
}
|
| 48 |
-
{%- endif %}
|
| 49 |
-
{% endmacro %}
|
| 50 |
-
|
| 51 |
{{ env.wgsl.resourceDeclarations }}
|
| 52 |
|
| 53 |
// com.microsoft.BiasAdd : Y = X + bias + skip
|
|
@@ -65,7 +38,7 @@ const HIDDEN: u32 = {{ hidden }}u;
|
|
| 65 |
{% if vec4 %}
|
| 66 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 67 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 68 |
-
{{ flat_index_2d() }}
|
| 69 |
let base = i * 4u;
|
| 70 |
let xv = vec4<f32>(f32(x[base]), f32(x[base + 1u]), f32(x[base + 2u]), f32(x[base + 3u]));
|
| 71 |
let sv = vec4<f32>(f32(skip[base]), f32(skip[base + 1u]), f32(skip[base + 2u]), f32(skip[base + 3u]));
|
|
@@ -80,6 +53,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
| 80 |
let sv = f32(skip[i]);
|
| 81 |
let v = xv + f32(bias[i % HIDDEN]) + sv;
|
| 82 |
y[i] = {{ scalar }}(v);
|
| 83 |
-
{{ flat_tail_close()
|
| 84 |
}
|
| 85 |
{% endif %}
|
|
|
|
| 12 |
for (var i = begin; i < end; i = i + 1u) {
|
| 13 |
{%- endmacro %}
|
| 14 |
{% macro flat_tail_close() %}
|
| 15 |
+
}{% endmacro %}
|
| 16 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 17 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 18 |
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 19 |
// per-axis workgroup fold width.
|
| 20 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
if ({{ name }} >= {{ bound }}) {
|
| 22 |
return;
|
| 23 |
+
}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
| 24 |
{{ env.wgsl.resourceDeclarations }}
|
| 25 |
|
| 26 |
// com.microsoft.BiasAdd : Y = X + bias + skip
|
|
|
|
| 38 |
{% if vec4 %}
|
| 39 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 40 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 41 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE) }}
|
| 42 |
let base = i * 4u;
|
| 43 |
let xv = vec4<f32>(f32(x[base]), f32(x[base + 1u]), f32(x[base + 2u]), f32(x[base + 3u]));
|
| 44 |
let sv = vec4<f32>(f32(skip[base]), f32(skip[base + 1u]), f32(skip[base + 2u]), f32(skip[base + 3u]));
|
|
|
|
| 53 |
let sv = f32(skip[i]);
|
| 54 |
let v = xv + f32(bias[i % HIDDEN]) + sv;
|
| 55 |
y[i] = {{ scalar }}(v);
|
| 56 |
+
{{ flat_tail_close() }}
|
| 57 |
}
|
| 58 |
{% endif %}
|
build/webgpu/manifest.json
CHANGED
|
@@ -12,7 +12,7 @@
|
|
| 12 |
"tunables": { "WORKGROUP_SIZE": { "default": 256 } },
|
| 13 |
"derive": { "scalar": "dtypes.T", "hidden": "dim(shapes.X, ranks.X - 1) if dim(shapes.X, ranks.X - 1) > 0 else 1" },
|
| 14 |
"when": ["ranks.X == 3", "ranks.skip == 3", "ranks.Y == 3", "sameShape(shapes.X, shapes.skip)", "sameShape(shapes.X, shapes.Y)", "f16Ok(dtypes.T)", "ranks.bias == 1", "dim(shapes.bias, 0) == dim(shapes.X, 2)"],
|
| 15 |
-
"bindings": { "x": { "arg": "X", "
|
| 16 |
"variants": [
|
| 17 |
{
|
| 18 |
"id": "vec4",
|
|
@@ -22,7 +22,7 @@
|
|
| 22 |
"passes": [
|
| 23 |
{
|
| 24 |
"id": "main",
|
| 25 |
-
"name": "BiasAdd.
|
| 26 |
"shader": "bias-add.wgsl.jinja",
|
| 27 |
"bindings": [
|
| 28 |
"x",
|
|
@@ -46,7 +46,7 @@
|
|
| 46 |
"passes": [
|
| 47 |
{
|
| 48 |
"id": "main",
|
| 49 |
-
"name": "BiasAdd.
|
| 50 |
"shader": "bias-add.wgsl.jinja",
|
| 51 |
"derive": { "itemsPerInvocation": 4 },
|
| 52 |
"bindings": [
|
|
|
|
| 12 |
"tunables": { "WORKGROUP_SIZE": { "default": 256 } },
|
| 13 |
"derive": { "scalar": "dtypes.T", "hidden": "dim(shapes.X, ranks.X - 1) if dim(shapes.X, ranks.X - 1) > 0 else 1" },
|
| 14 |
"when": ["ranks.X == 3", "ranks.skip == 3", "ranks.Y == 3", "sameShape(shapes.X, shapes.skip)", "sameShape(shapes.X, shapes.Y)", "f16Ok(dtypes.T)", "ranks.bias == 1", "dim(shapes.bias, 0) == dim(shapes.X, 2)"],
|
| 15 |
+
"bindings": { "x": { "arg": "X", "elementType": "$scalar" } },
|
| 16 |
"variants": [
|
| 17 |
{
|
| 18 |
"id": "vec4",
|
|
|
|
| 22 |
"passes": [
|
| 23 |
{
|
| 24 |
"id": "main",
|
| 25 |
+
"name": "BiasAdd.Vec4",
|
| 26 |
"shader": "bias-add.wgsl.jinja",
|
| 27 |
"bindings": [
|
| 28 |
"x",
|
|
|
|
| 46 |
"passes": [
|
| 47 |
{
|
| 48 |
"id": "main",
|
| 49 |
+
"name": "BiasAdd.Scalar",
|
| 50 |
"shader": "bias-add.wgsl.jinja",
|
| 51 |
"derive": { "itemsPerInvocation": 4 },
|
| 52 |
"bindings": [
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.BiasAdd",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
@@ -8,14 +8,14 @@
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "RdL1BQwfXYch113DNTUHHlTuPCDk1Pp6wKa2sV9SeY4=",
|
| 11 |
-
"bias-add.wgsl.jinja": "
|
| 12 |
-
"manifest.json": "
|
| 13 |
"test.json": "yECzeltqwOXVNMeShcml5Bz3B2Amx3fmBX47DQCPG70="
|
| 14 |
}
|
| 15 |
},
|
| 16 |
-
"provenance": { "kernel": { "sha": "
|
| 17 |
"webgpu": {
|
| 18 |
-
"manifestSpec": "2.
|
| 19 |
"variants": { "vec4": ["bias-add.wgsl.jinja"], "scalar": ["bias-add.wgsl.jinja"] }
|
| 20 |
}
|
| 21 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.BiasAdd",
|
| 3 |
+
"id": "_com_microsoft_biasadd_webgpu_a85a8a8",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "RdL1BQwfXYch113DNTUHHlTuPCDk1Pp6wKa2sV9SeY4=",
|
| 11 |
+
"bias-add.wgsl.jinja": "SX9BxzKhgcHf/An4OvUU/YM9yZ5nKiZClOQZiy7L1Kw=",
|
| 12 |
+
"manifest.json": "kmHUwwH3j3g8PsJzhnbKF0tku6/rh3LWTAmDkHAyQGA=",
|
| 13 |
"test.json": "yECzeltqwOXVNMeShcml5Bz3B2Amx3fmBX47DQCPG70="
|
| 14 |
}
|
| 15 |
},
|
| 16 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 17 |
"webgpu": {
|
| 18 |
+
"manifestSpec": "2.1",
|
| 19 |
"variants": { "vec4": ["bias-add.wgsl.jinja"], "scalar": ["bias-add.wgsl.jinja"] }
|
| 20 |
}
|
| 21 |
}
|