Xenova HF Staff commited on
Commit
be4f8dc
·
verified ·
1 Parent(s): 61c5e0b

sync 6fdf6301e2bb

Browse files
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 + tuning 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.2
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
- {% endmacro %}
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
- {% endif %}
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", "buffer": "read-only-storage", "elementType": "$scalar" } },
16
  "variants": [
17
  {
18
  "id": "vec4",
@@ -22,7 +22,7 @@
22
  "passes": [
23
  {
24
  "id": "main",
25
- "name": "BiasAdd.vec4",
26
  "shader": "bias-add.wgsl.jinja",
27
  "bindings": [
28
  "x",
@@ -46,7 +46,7 @@
46
  "passes": [
47
  {
48
  "id": "main",
49
- "name": "BiasAdd.scalar",
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": "_com_microsoft_biasadd_webgpu_7687b51",
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": "w8g/0fT1IBapNvGjVIocSxND3JpDsMvxJQYZHQJOAXw=",
12
- "manifest.json": "iwu1wkLNtggvjZU/fzw2dq+3bcm+q/KfLIYAUpManhY=",
13
  "test.json": "yECzeltqwOXVNMeShcml5Bz3B2Amx3fmBX47DQCPG70="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
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
  }