Xenova HF Staff commited on
Commit
269fe3e
·
verified ·
1 Parent(s): 0500624

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -47,13 +47,13 @@ Default values (overridable per request):
47
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
48
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
49
  - [`test.json`](build/webgpu/test.json) — correctness cases
50
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
51
  - [`quick-gelu.wgsl.jinja`](build/webgpu/quick-gelu.wgsl.jinja)
52
 
53
  ## Use with `@huggingface/kernels`
54
 
55
  ```sh
56
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
57
  ```
58
 
59
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
47
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
48
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
49
  - [`test.json`](build/webgpu/test.json) — correctness cases
50
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
51
  - [`quick-gelu.wgsl.jinja`](build/webgpu/quick-gelu.wgsl.jinja)
52
 
53
  ## Use with `@huggingface/kernels`
54
 
55
  ```sh
56
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
57
  ```
58
 
59
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/manifest.json CHANGED
@@ -11,7 +11,9 @@
11
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
12
  "workgroupOk": "tunables.WORKGROUP_SIZE > 0 and tunables.WORKGROUP_SIZE <= deviceWorkgroupCap",
13
  "baseOk": "workgroupOk and numel(shapes.X) == numel(shapes.Y) and f16Ok(dtypes.T)",
14
- "vec4Ok": "numel(shapes.X) > 0 and numel(shapes.X) % 4 == 0"
 
 
15
  },
16
  "when": ["baseOk"],
17
  "variants": [
@@ -19,18 +21,12 @@
19
  "id": "vec4",
20
  "priority": 20,
21
  "when": ["vec4Ok"],
22
- "derive": {
23
- "scalar": "dtypes.T",
24
- "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
25
- "vec4": true,
26
- "vec4Tail": false
27
- },
28
  "passes": [
29
  {
30
  "id": "main",
31
- "name": "QuickGelu.vec4",
32
  "shader": "quick-gelu.wgsl.jinja",
33
- "derive": { "alpha": "attrs.alpha" },
34
  "bindings": [
35
  { "arg": "X", "name": "x", "elementType": "$vectorScalar" },
36
  { "arg": "Y", "name": "y", "elementType": "$vectorScalar" },
@@ -47,14 +43,12 @@
47
  {
48
  "id": "vec4_tail",
49
  "priority": 10,
50
- "when": ["numel(shapes.X) > 0"],
51
- "derive": { "scalar": "dtypes.T", "vec4": false, "vec4Tail": true },
52
  "passes": [
53
  {
54
  "id": "main",
55
- "name": "QuickGelu.vec4Tail",
56
  "shader": "quick-gelu.wgsl.jinja",
57
- "derive": { "alpha": "attrs.alpha" },
58
  "bindings": [
59
  { "arg": "X", "name": "x", "elementType": "$scalar" },
60
  { "arg": "Y", "name": "y", "elementType": "$scalar" },
@@ -67,30 +61,6 @@
67
  }
68
  }
69
  ]
70
- },
71
- {
72
- "id": "scalar",
73
- "priority": 0,
74
- "when": ["true"],
75
- "derive": { "scalar": "dtypes.T", "vec4": false, "vec4Tail": false },
76
- "passes": [
77
- {
78
- "id": "main",
79
- "name": "QuickGelu.scalar",
80
- "shader": "quick-gelu.wgsl.jinja",
81
- "derive": { "alpha": "attrs.alpha" },
82
- "bindings": [
83
- { "arg": "X", "name": "x", "elementType": "$scalar" },
84
- { "arg": "Y", "name": "y", "elementType": "$scalar" },
85
- { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.X)" }] }
86
- ],
87
- "dispatch": {
88
- "x": "min(ceilDiv((numel(shapes.X)), (tunables.WORKGROUP_SIZE)), 65535)",
89
- "y": "ceilDiv(ceilDiv((numel(shapes.X)), (tunables.WORKGROUP_SIZE)), 65535)",
90
- "z": 1
91
- }
92
- }
93
- ]
94
  }
95
  ]
96
  }
 
11
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
12
  "workgroupOk": "tunables.WORKGROUP_SIZE > 0 and tunables.WORKGROUP_SIZE <= deviceWorkgroupCap",
13
  "baseOk": "workgroupOk and numel(shapes.X) == numel(shapes.Y) and f16Ok(dtypes.T)",
14
+ "vec4Ok": "numel(shapes.X) > 0 and numel(shapes.X) % 4 == 0",
15
+ "scalar": "dtypes.T",
16
+ "alpha": "attrs.alpha"
17
  },
18
  "when": ["baseOk"],
19
  "variants": [
 
21
  "id": "vec4",
22
  "priority": 20,
23
  "when": ["vec4Ok"],
24
+ "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"", "vec4": true, "vec4Tail": false },
 
 
 
 
 
25
  "passes": [
26
  {
27
  "id": "main",
28
+ "name": "QuickGelu.Vec4",
29
  "shader": "quick-gelu.wgsl.jinja",
 
30
  "bindings": [
31
  { "arg": "X", "name": "x", "elementType": "$vectorScalar" },
32
  { "arg": "Y", "name": "y", "elementType": "$vectorScalar" },
 
43
  {
44
  "id": "vec4_tail",
45
  "priority": 10,
46
+ "derive": { "vec4": false, "vec4Tail": "numel(shapes.X) > 0" },
 
47
  "passes": [
48
  {
49
  "id": "main",
50
+ "name": "QuickGelu.Vec4Tail",
51
  "shader": "quick-gelu.wgsl.jinja",
 
52
  "bindings": [
53
  { "arg": "X", "name": "x", "elementType": "$scalar" },
54
  { "arg": "Y", "name": "y", "elementType": "$scalar" },
 
61
  }
62
  }
63
  ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
  }
65
  ]
66
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.QuickGelu",
3
- "id": "_com_microsoft_quickgelu_webgpu_57d870f",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,18 +8,14 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "UkzjTMKrQBdgOKbmqRYDd53jtc70n2SBbpJvjok4i60=",
11
- "manifest.json": "BwEkPwdzXku28VYZ+dMC7eKDL4HzAlj6H2Iwx61Rmw8=",
12
- "quick-gelu.wgsl.jinja": "/+KTUHzwuZwWV4Q2df6YEdgv2b5wu6EBdI1eZUCUj10=",
13
  "test.json": "yMUdwnbXzdJruLruAgEJ6bcftDu0hE1av8ovE2UsXQs="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
- "variants": {
20
- "vec4": ["quick-gelu.wgsl.jinja"],
21
- "vec4_tail": ["quick-gelu.wgsl.jinja"],
22
- "scalar": ["quick-gelu.wgsl.jinja"]
23
- }
24
  }
25
  }
 
1
  {
2
  "name": "com.microsoft.QuickGelu",
3
+ "id": "_com_microsoft_quickgelu_webgpu_f22c71c",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "UkzjTMKrQBdgOKbmqRYDd53jtc70n2SBbpJvjok4i60=",
11
+ "manifest.json": "ZWAG7Y3eLNxEgaO8Em5Bd+x1k1yJ8qZ34hXYKyPxvOE=",
12
+ "quick-gelu.wgsl.jinja": "nRDogpODhEJqpNDO+dS/wROsmB29cgLabNzUW2VqfNo=",
13
  "test.json": "yMUdwnbXzdJruLruAgEJ6bcftDu0hE1av8ovE2UsXQs="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
  "webgpu": {
18
+ "manifestSpec": "2.1",
19
+ "variants": { "vec4": ["quick-gelu.wgsl.jinja"], "vec4_tail": ["quick-gelu.wgsl.jinja"] }
 
 
 
 
20
  }
21
  }
build/webgpu/quick-gelu.wgsl.jinja CHANGED
@@ -1,36 +1,11 @@
1
- {% macro flat_index_2d(name="i", bound="params.count", guardInline=false, note="dispatch-limit") %}
2
- {% if note == "dispatch-limit" %}
3
- // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
- // per-axis workgroup fold width (outputs > 16.7M elements).
5
- {% elif note == "limit" %}
6
- // 2D-folded flat index: gid.y carries the high bits past the dispatch's
7
- // per-axis workgroup fold width.
8
- {% elif note == "device-axis" %}
9
- // The flat dispatch is folded across x/y at a fixed per-axis workgroup
10
- // width; gid.y carries the high portion of the output index.
11
- {% elif note == "vec4-limit" %}
12
- // 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
13
- // per-axis workgroup fold width (the dispatch caps x and spills into y).
14
- {% elif note == "element-limit" %}
15
- // 2D-folded flat element index: gid.y carries the high bits past the
16
- // dispatch's per-axis workgroup fold width.
17
- {% elif note == "dispatch" %}
18
  // 2D-folded flat index: gid.y carries the high bits past the dispatch's
19
  // per-axis workgroup fold width.
20
- {% endif %}
21
- {% if bound == "" %}
22
- let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
23
- {%- elif guardInline %}
24
- let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
25
- if ({{ name }} >= {{ bound }}) { return; }
26
- {%- else %}
27
- let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
28
  if ({{ name }} >= {{ bound }}) {
29
  return;
30
- }
31
- {%- endif %}
32
- {% endmacro %}
33
-
34
  {{ env.wgsl.resourceDeclarations }}
35
 
36
  // com.microsoft.QuickGelu : Y = X * sigmoid(alpha * X)
@@ -58,7 +33,7 @@ fn quick_gelu(v: f32) -> f32 {
58
 
59
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
60
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
61
- {{ flat_index_2d() }}
62
  {% if vec4Tail %}
63
  let base = i * 4u;
64
  {% for lane in range(4) %}
 
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
  // com.microsoft.QuickGelu : Y = X * sigmoid(alpha * X)
 
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
  {% if vec4Tail %}
38
  let base = i * 4u;
39
  {% for lane in range(4) %}