Xenova HF Staff commited on
Commit
2cf8441
·
verified ·
1 Parent(s): 53d66ee

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -39,13 +39,13 @@ See the [ONNX Runtime `Gelu` contrib-operator spec](https://github.com/microsoft
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 + tuning cases
43
  - [`elementwise-bias-gelu.wgsl.jinja`](build/webgpu/elementwise-bias-gelu.wgsl.jinja)
44
 
45
  ## Use with `@huggingface/kernels`
46
 
47
  ```sh
48
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
49
  ```
50
 
51
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
@@ -59,5 +59,5 @@ Replace each `*Data` placeholder with a typed array containing the corresponding
59
  import { getKernel } from "@huggingface/kernels";
60
 
61
  const kernel = await getKernel("webgpu-kernels/com.microsoft.Gelu", { version: 1 });
62
- const { Y } = await kernel({ X: { data: XData, shape: [] } });
63
  ```
 
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
  - [`elementwise-bias-gelu.wgsl.jinja`](build/webgpu/elementwise-bias-gelu.wgsl.jinja)
44
 
45
  ## Use with `@huggingface/kernels`
46
 
47
  ```sh
48
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
49
  ```
50
 
51
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
59
  import { getKernel } from "@huggingface/kernels";
60
 
61
  const kernel = await getKernel("webgpu-kernels/com.microsoft.Gelu", { version: 1 });
62
+ const { Y } = await kernel({ X: { data: XData, shape: [5] } });
63
  ```
build/webgpu/elementwise-bias-gelu.wgsl.jinja CHANGED
@@ -1,3 +1,11 @@
 
 
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
  {% set wg = workgroupSize if workgroupSize is defined else tunables.WORKGROUP_SIZE %}
3
 
@@ -20,17 +28,14 @@ fn erf_approx(x: f32) -> f32 {
20
  let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t * exp(-(ax * ax));
21
  return sign * y;
22
  }
 
23
  fn gelu_value(v: f32) -> f32 {
24
  return 0.5 * v * (1.0 + erf_approx(v * 0.7071067811865476));
25
  }
 
26
  @compute @workgroup_size({{ wg }})
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 (outputs > 16.7M elements).
30
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wg }}u;
31
- if (i >= params.count) {
32
- return;
33
- }
34
  {% if vec4 %}
35
  let xv = vec4<f32>(x[i]);
36
  let v = xv;
 
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
  {% set wg = workgroupSize if workgroupSize is defined else tunables.WORKGROUP_SIZE %}
11
 
 
28
  let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t * exp(-(ax * ax));
29
  return sign * y;
30
  }
31
+
32
  fn gelu_value(v: f32) -> f32 {
33
  return 0.5 * v * (1.0 + erf_approx(v * 0.7071067811865476));
34
  }
35
+
36
  @compute @workgroup_size({{ wg }})
37
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
38
+ {{ flat_index_2d(wg) }}
 
 
 
 
 
39
  {% if vec4 %}
40
  let xv = vec4<f32>(x[i]);
41
  let v = xv;
build/webgpu/manifest.json CHANGED
@@ -6,7 +6,7 @@
6
  "outputs": { "Y": { "dtype": "T", "rank": "ranks.X", "shape": "shapes.X" } },
7
  "typeConstraints": { "T": ["float32", "float16"] },
8
  "tunables": { "WORKGROUP_SIZE": { "default": 256 } },
9
- "derive": { "scalar": "dtypes.T", "approximate": "\"erf\"", "vec4Tail": false, "hasBias": false },
10
  "when": ["numel(shapes.X) == numel(shapes.Y)", "f16Ok(dtypes.T)"],
11
  "variants": [
12
  {
@@ -17,7 +17,7 @@
17
  "passes": [
18
  {
19
  "id": "main",
20
- "name": "Gelu.vec4",
21
  "shader": "elementwise-bias-gelu.wgsl.jinja",
22
  "bindings": [
23
  { "arg": "X", "name": "x", "elementType": "$vectorScalar" },
@@ -39,7 +39,7 @@
39
  "passes": [
40
  {
41
  "id": "main",
42
- "name": "Gelu.scalar",
43
  "shader": "elementwise-bias-gelu.wgsl.jinja",
44
  "bindings": [
45
  { "arg": "X", "name": "x", "elementType": "$scalar" },
 
6
  "outputs": { "Y": { "dtype": "T", "rank": "ranks.X", "shape": "shapes.X" } },
7
  "typeConstraints": { "T": ["float32", "float16"] },
8
  "tunables": { "WORKGROUP_SIZE": { "default": 256 } },
9
+ "derive": { "scalar": "dtypes.T", "vec4Tail": false },
10
  "when": ["numel(shapes.X) == numel(shapes.Y)", "f16Ok(dtypes.T)"],
11
  "variants": [
12
  {
 
17
  "passes": [
18
  {
19
  "id": "main",
20
+ "name": "Gelu.Vec4",
21
  "shader": "elementwise-bias-gelu.wgsl.jinja",
22
  "bindings": [
23
  { "arg": "X", "name": "x", "elementType": "$vectorScalar" },
 
39
  "passes": [
40
  {
41
  "id": "main",
42
+ "name": "Gelu.Scalar",
43
  "shader": "elementwise-bias-gelu.wgsl.jinja",
44
  "bindings": [
45
  { "arg": "X", "name": "x", "elementType": "$scalar" },
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.Gelu",
3
- "id": "_com_microsoft_gelu_webgpu_9075789",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,14 +8,14 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "zb+BUPCMklZQithfY94RzWttjR/SZK0PIWgZjexuaxs=",
11
- "elementwise-bias-gelu.wgsl.jinja": "K4LR9m7TibnO/vuJpCLk9+CTk1slr+E/XCGlmgH7hww=",
12
- "manifest.json": "R8nYWq7dFuIme6gF1PUJWYKRv8u6u67dxzGzsZFLs38=",
13
  "test.json": "K4VvkDT8J9wTZh+uYL3GSIAzvOw0PuV04RAIyhxmgxM="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
  "variants": { "vec4": ["elementwise-bias-gelu.wgsl.jinja"], "scalar": ["elementwise-bias-gelu.wgsl.jinja"] }
20
  }
21
  }
 
1
  {
2
  "name": "com.microsoft.Gelu",
3
+ "id": "_com_microsoft_gelu_webgpu_7e75dcc",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "zb+BUPCMklZQithfY94RzWttjR/SZK0PIWgZjexuaxs=",
11
+ "elementwise-bias-gelu.wgsl.jinja": "BDGBTzz/SKEd8tbwt6VCzjky2+w28H9UZ6WSGEbeZSQ=",
12
+ "manifest.json": "Vx/CPPXNBVsauc5qnG9IeRU55D0gQVbGJtk/S40peXg=",
13
  "test.json": "K4VvkDT8J9wTZh+uYL3GSIAzvOw0PuV04RAIyhxmgxM="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
  "webgpu": {
18
+ "manifestSpec": "2.1",
19
  "variants": { "vec4": ["elementwise-bias-gelu.wgsl.jinja"], "scalar": ["elementwise-bias-gelu.wgsl.jinja"] }
20
  }
21
  }