Xenova HF Staff commited on
Commit
5def70a
·
verified ·
1 Parent(s): 5e78980

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -41,13 +41,13 @@ See the [ONNX Runtime `GatedAdd` contrib-operator spec](https://github.com/micro
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
  - [`gated-add.wgsl.jinja`](build/webgpu/gated-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
  - [`gated-add.wgsl.jinja`](build/webgpu/gated-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/gated-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.GatedAdd : output = X + round_to_T(Y * gate)
@@ -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
  // A vec4 group is four consecutive channels of one row: HIDDEN % 4 == 0 stops
70
  // it from ever straddling two rows, so the whole group shares one gate value.
71
  let g = vec4<{{ scalar }}>(gate[i * 4u / HIDDEN]);
@@ -75,6 +48,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
75
  {{ flat_tail_open() }}
76
  let g = gate[i / HIDDEN];
77
  output[i] = x[i] + fma(y[i], g, {{ scalar }}(0.0));
78
- {{ flat_tail_close() -}}
79
  }
80
  {% 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.GatedAdd : output = X + round_to_T(Y * gate)
 
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
  // A vec4 group is four consecutive channels of one row: HIDDEN % 4 == 0 stops
43
  // it from ever straddling two rows, so the whole group shares one gate value.
44
  let g = vec4<{{ scalar }}>(gate[i * 4u / HIDDEN]);
 
48
  {{ flat_tail_open() }}
49
  let g = gate[i / HIDDEN];
50
  output[i] = x[i] + fma(y[i], g, {{ scalar }}(0.0));
51
+ {{ flat_tail_close() }}
52
  }
53
  {% endif %}
build/webgpu/manifest.json CHANGED
@@ -23,7 +23,7 @@
23
  "passes": [
24
  {
25
  "id": "main",
26
- "name": "GatedAdd.vec4",
27
  "shader": "gated-add.wgsl.jinja",
28
  "bindings": [
29
  { "arg": "X", "name": "x", "elementType": "$vectorScalar" },
@@ -47,7 +47,7 @@
47
  "passes": [
48
  {
49
  "id": "main",
50
- "name": "GatedAdd.scalar",
51
  "shader": "gated-add.wgsl.jinja",
52
  "derive": { "itemsPerInvocation": 4 },
53
  "bindings": [
 
23
  "passes": [
24
  {
25
  "id": "main",
26
+ "name": "GatedAdd.Vec4",
27
  "shader": "gated-add.wgsl.jinja",
28
  "bindings": [
29
  { "arg": "X", "name": "x", "elementType": "$vectorScalar" },
 
47
  "passes": [
48
  {
49
  "id": "main",
50
+ "name": "GatedAdd.Scalar",
51
  "shader": "gated-add.wgsl.jinja",
52
  "derive": { "itemsPerInvocation": 4 },
53
  "bindings": [
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.GatedAdd",
3
- "id": "_com_microsoft_gatedadd_webgpu_913bb83",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,14 +8,14 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "f/TmVS5V2Nom2d9yiZ8HU0gLA+HtwVePXaGO4EmsmoY=",
11
- "gated-add.wgsl.jinja": "6xWRekBrMasrrGGXBN6cHj1NeycAppVTPGz1QPRDnUQ=",
12
- "manifest.json": "a3y5tgF4lud3fMIhIIyomyBnUdsTqCeCaWuGJOSdW1U=",
13
- "test.json": "CgQChPZDm5zqD3shwXcy5E57YHqJlw8wSwSyTQeEvN4="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
  "variants": { "vec4": ["gated-add.wgsl.jinja"], "scalar": ["gated-add.wgsl.jinja"] }
20
  }
21
  }
 
1
  {
2
  "name": "com.microsoft.GatedAdd",
3
+ "id": "_com_microsoft_gatedadd_webgpu_34aae76",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "f/TmVS5V2Nom2d9yiZ8HU0gLA+HtwVePXaGO4EmsmoY=",
11
+ "gated-add.wgsl.jinja": "jZnIxRERHOaeJx2usJ13f9c7XLMB2bVMMIfqZ1DMYu8=",
12
+ "manifest.json": "PCK9cwNdN2BqtsElsaqlUDXwrON3HQA3nxwc8lmw2LM=",
13
+ "test.json": "Mu0iaJZ9Ygpmyu+1cY1bU6ytPO2Q0eYGf401O3L6o6w="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
  "webgpu": {
18
+ "manifestSpec": "2.1",
19
  "variants": { "vec4": ["gated-add.wgsl.jinja"], "scalar": ["gated-add.wgsl.jinja"] }
20
  }
21
  }
build/webgpu/test.json CHANGED
@@ -179,6 +179,84 @@
179
  "gate": { "dtype": "float16", "shape": [4, 1], "data": { "kind": "values", "values": [0.75, -1.25, 3.0, 0.5] } }
180
  },
181
  "outputs": { "output": { "dtype": "float16", "shape": [4, 5], "tolerance": 0.0005, "relTolerance": 0.002 } }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
182
  }
183
  ]
184
  }
 
179
  "gate": { "dtype": "float16", "shape": [4, 1], "data": { "kind": "values", "values": [0.75, -1.25, 3.0, 0.5] } }
180
  },
181
  "outputs": { "output": { "dtype": "float16", "shape": [4, 5], "tolerance": 0.0005, "relTolerance": 0.002 } }
182
+ },
183
+ {
184
+ "name": "ort_gated_add_rank3_hidden7_f32",
185
+ "inputs": {
186
+ "X": {
187
+ "dtype": "float32",
188
+ "shape": [2, 3, 7],
189
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29 }
190
+ },
191
+ "Y": {
192
+ "dtype": "float32",
193
+ "shape": [2, 3, 7],
194
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.07 }
195
+ },
196
+ "gate": {
197
+ "dtype": "float32",
198
+ "shape": [2, 3, 1],
199
+ "data": { "kind": "fillFloat32", "sinStep": 0.41, "cosStep": 0.17, "offset": 0.75 }
200
+ }
201
+ },
202
+ "outputs": {
203
+ "output": { "dtype": "float32", "shape": [2, 3, 7], "tolerance": 0.000001, "relTolerance": 0.000001 }
204
+ }
205
+ },
206
+ {
207
+ "name": "ort_gated_add_rank1_hidden7_f32",
208
+ "inputs": {
209
+ "X": { "dtype": "float32", "shape": [7], "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29 } },
210
+ "Y": { "dtype": "float32", "shape": [7], "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.07 } },
211
+ "gate": {
212
+ "dtype": "float32",
213
+ "shape": [1],
214
+ "data": { "kind": "fillFloat32", "sinStep": 0.41, "cosStep": 0.17, "offset": 0.75 }
215
+ }
216
+ },
217
+ "outputs": { "output": { "dtype": "float32", "shape": [7], "tolerance": 0.000001, "relTolerance": 0.000001 } }
218
+ },
219
+ {
220
+ "name": "ort_gated_add_empty_outer_dim_f32",
221
+ "inputs": {
222
+ "X": {
223
+ "dtype": "float32",
224
+ "shape": [0, 7],
225
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29 }
226
+ },
227
+ "Y": {
228
+ "dtype": "float32",
229
+ "shape": [0, 7],
230
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.07 }
231
+ },
232
+ "gate": {
233
+ "dtype": "float32",
234
+ "shape": [0, 1],
235
+ "data": { "kind": "fillFloat32", "sinStep": 0.41, "cosStep": 0.17, "offset": 0.75 }
236
+ }
237
+ },
238
+ "outputs": { "output": { "dtype": "float32", "shape": [0, 7], "tolerance": 0.000001, "relTolerance": 0.000001 } }
239
+ },
240
+ {
241
+ "name": "ort_gated_add_f16_hidden2048",
242
+ "inputs": {
243
+ "X": {
244
+ "dtype": "float16",
245
+ "shape": [1, 4, 2048],
246
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29 }
247
+ },
248
+ "Y": {
249
+ "dtype": "float16",
250
+ "shape": [1, 4, 2048],
251
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.07 }
252
+ },
253
+ "gate": {
254
+ "dtype": "float16",
255
+ "shape": [1, 4, 1],
256
+ "data": { "kind": "fillFloat32", "sinStep": 0.41, "cosStep": 0.17, "offset": 0.75 }
257
+ }
258
+ },
259
+ "outputs": { "output": { "dtype": "float16", "shape": [1, 4, 2048], "tolerance": 0.0005, "relTolerance": 0.002 } }
260
  }
261
  ]
262
  }