Xenova HF Staff commited on
Commit
fe69102
·
verified ·
1 Parent(s): 02e3be6

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -40,13 +40,13 @@ See the [ONNX Runtime `FastGelu` contrib-operator spec](https://github.com/micro
40
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
41
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
42
  - [`test.json`](build/webgpu/test.json) — correctness cases
43
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
44
  - [`elementwise-bias-gelu.wgsl.jinja`](build/webgpu/elementwise-bias-gelu.wgsl.jinja)
45
 
46
  ## Use with `@huggingface/kernels`
47
 
48
  ```sh
49
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
50
  ```
51
 
52
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
40
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
41
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
42
  - [`test.json`](build/webgpu/test.json) — correctness cases
43
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
44
  - [`elementwise-bias-gelu.wgsl.jinja`](build/webgpu/elementwise-bias-gelu.wgsl.jinja)
45
 
46
  ## Use with `@huggingface/kernels`
47
 
48
  ```sh
49
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
50
  ```
51
 
52
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
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
 
@@ -7,26 +15,27 @@
7
  // `vec4Tail` instead uses scalar bindings with four guarded lanes.
8
  // GELU uses the tanh approximation below, with its input clamped in the tails.
9
  fn tanh_safe(x: f32) -> f32 {
 
 
10
  if (x > 10.0) { return 1.0; }
11
  if (x < -10.0) { return -1.0; }
 
 
 
12
  return tanh(x);
13
  }
 
14
  fn gelu_value(v: f32) -> f32 {
15
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
16
  }
17
- {% if hasBias %}
18
 
 
19
  const HIDDEN: u32 = {{ hidden | default(0) }}u;
20
 
21
  {% endif %}
22
  @compute @workgroup_size({{ wg }})
23
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
24
- // 2D-folded flat index: gid.y carries the high bits past the
25
- // per-axis dispatch fold width (outputs > 16.7M elements).
26
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wg }}u;
27
- if (i >= params.count) {
28
- return;
29
- }
30
  {% if vec4Tail %}
31
  let base = i * 4u;
32
  {% 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
  {% set wg = workgroupSize if workgroupSize is defined else tunables.WORKGROUP_SIZE %}
11
 
 
15
  // `vec4Tail` instead uses scalar bindings with four guarded lanes.
16
  // GELU uses the tanh approximation below, with its input clamped in the tails.
17
  fn tanh_safe(x: f32) -> f32 {
18
+ // tanh rounds to its saturated value for these tails in f32. Return that
19
+ // value directly, including for infinite input, before invoking the builtin.
20
  if (x > 10.0) { return 1.0; }
21
  if (x < -10.0) { return -1.0; }
22
+ // For tiny |x|, return x directly to preserve its sign and magnitude without
23
+ // relying on backend-specific builtin behavior near zero.
24
+ if (x > -1.0e-4 && x < 1.0e-4) { return x; }
25
  return tanh(x);
26
  }
27
+
28
  fn gelu_value(v: f32) -> f32 {
29
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
30
  }
 
31
 
32
+ {% if hasBias %}
33
  const HIDDEN: u32 = {{ hidden | default(0) }}u;
34
 
35
  {% endif %}
36
  @compute @workgroup_size({{ wg }})
37
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
38
+ {{ flat_index_2d(wg) }}
 
 
 
 
 
39
  {% if vec4Tail %}
40
  let base = i * 4u;
41
  {% for lane in range(4) %}
build/webgpu/manifest.json CHANGED
@@ -16,21 +16,16 @@
16
  "biasOk": "present.bias and ranks.bias == 1 and dim(shapes.bias, 0) == dim(shapes.X, ranks.X - 1)",
17
  "noBiasOk": "not present.bias",
18
  "vec4Ok": "numel(shapes.X) > 0 and numel(shapes.X) % 4 == 0 and dim(shapes.X, ranks.X - 1) % 4 == 0",
19
- "scalar": "dtypes.T",
20
- "approximate": "\"tanh\""
21
  },
22
  "when": ["baseOk"],
23
  "bindings": {
24
- "x": { "arg": "X", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
25
- "y": { "arg": "Y", "buffer": "storage", "elementType": "$vectorScalar" },
26
- "params": { "buffer": "uniform", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.X) / 4" }] },
27
- "x_2": { "arg": "X", "name": "x", "buffer": "read-only-storage", "elementType": "$scalar" },
28
- "y_2": { "arg": "Y", "name": "y", "buffer": "storage", "elementType": "$scalar" },
29
- "params_2": {
30
- "name": "params",
31
- "buffer": "uniform",
32
- "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.X)" }]
33
- }
34
  },
35
  "variants": [
36
  {
@@ -47,7 +42,7 @@
47
  "passes": [
48
  {
49
  "id": "main",
50
- "name": "FastGelu.vec4Bias",
51
  "shader": "elementwise-bias-gelu.wgsl.jinja",
52
  "bindings": ["x", { "arg": "bias", "elementType": "$scalar", "length": "$hidden" }, "y", "params"],
53
  "dispatch": {
@@ -66,7 +61,7 @@
66
  "passes": [
67
  {
68
  "id": "main",
69
- "name": "FastGelu.vec4",
70
  "shader": "elementwise-bias-gelu.wgsl.jinja",
71
  "bindings": ["x", "y", "params"],
72
  "dispatch": {
@@ -80,19 +75,19 @@
80
  {
81
  "id": "vec4_tail_bias",
82
  "priority": 20,
83
- "when": ["biasOk", "numel(shapes.X) > 0"],
84
  "derive": {
85
  "vec4": false,
86
- "vec4Tail": true,
87
  "hasBias": true,
88
  "hidden": "dim(shapes.X, ranks.X - 1) if dim(shapes.X, ranks.X - 1) > 0 else 1"
89
  },
90
  "passes": [
91
  {
92
  "id": "main",
93
- "name": "FastGelu.vec4TailBias",
94
  "shader": "elementwise-bias-gelu.wgsl.jinja",
95
- "bindings": ["x_2", "bias", "y_2", "params_2"],
96
  "dispatch": {
97
  "x": "min(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
98
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
@@ -104,14 +99,14 @@
104
  {
105
  "id": "vec4_tail_no_bias",
106
  "priority": 15,
107
- "when": ["noBiasOk", "numel(shapes.X) > 0"],
108
- "derive": { "vec4": false, "vec4Tail": true, "hasBias": false },
109
  "passes": [
110
  {
111
  "id": "main",
112
- "name": "FastGelu.vec4Tail",
113
  "shader": "elementwise-bias-gelu.wgsl.jinja",
114
- "bindings": ["x_2", "y_2", "params_2"],
115
  "dispatch": {
116
  "x": "min(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
117
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
@@ -119,49 +114,6 @@
119
  }
120
  }
121
  ]
122
- },
123
- {
124
- "id": "scalar_bias",
125
- "priority": 10,
126
- "when": ["biasOk", "true"],
127
- "derive": {
128
- "vec4": false,
129
- "vec4Tail": false,
130
- "hasBias": true,
131
- "hidden": "dim(shapes.X, ranks.X - 1) if dim(shapes.X, ranks.X - 1) > 0 else 1"
132
- },
133
- "passes": [
134
- {
135
- "id": "main",
136
- "name": "FastGelu.scalarBias",
137
- "shader": "elementwise-bias-gelu.wgsl.jinja",
138
- "bindings": ["x_2", "bias", "y_2", "params_2"],
139
- "dispatch": {
140
- "x": "min(ceilDiv((numel(shapes.X)), (workgroupSize)), 65535)",
141
- "y": "ceilDiv(ceilDiv((numel(shapes.X)), (workgroupSize)), 65535)",
142
- "z": 1
143
- }
144
- }
145
- ]
146
- },
147
- {
148
- "id": "scalar_no_bias",
149
- "priority": 0,
150
- "when": ["noBiasOk", "true"],
151
- "derive": { "vec4": false, "vec4Tail": false, "hasBias": false },
152
- "passes": [
153
- {
154
- "id": "main",
155
- "name": "FastGelu.scalar",
156
- "shader": "elementwise-bias-gelu.wgsl.jinja",
157
- "bindings": ["x_2", "y_2", "params_2"],
158
- "dispatch": {
159
- "x": "min(ceilDiv((numel(shapes.X)), (workgroupSize)), 65535)",
160
- "y": "ceilDiv(ceilDiv((numel(shapes.X)), (workgroupSize)), 65535)",
161
- "z": 1
162
- }
163
- }
164
- ]
165
  }
166
  ]
167
  }
 
16
  "biasOk": "present.bias and ranks.bias == 1 and dim(shapes.bias, 0) == dim(shapes.X, ranks.X - 1)",
17
  "noBiasOk": "not present.bias",
18
  "vec4Ok": "numel(shapes.X) > 0 and numel(shapes.X) % 4 == 0 and dim(shapes.X, ranks.X - 1) % 4 == 0",
19
+ "scalar": "dtypes.T"
 
20
  },
21
  "when": ["baseOk"],
22
  "bindings": {
23
+ "x": { "arg": "X", "elementType": "$vectorScalar" },
24
+ "y": { "arg": "Y", "elementType": "$vectorScalar" },
25
+ "params": { "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.X) / 4" }] },
26
+ "x_x": { "arg": "X", "name": "x", "elementType": "$scalar" },
27
+ "y_y": { "arg": "Y", "name": "y", "elementType": "$scalar" },
28
+ "params_main": { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.X)" }] }
 
 
 
 
29
  },
30
  "variants": [
31
  {
 
42
  "passes": [
43
  {
44
  "id": "main",
45
+ "name": "FastGelu.Vec4Bias",
46
  "shader": "elementwise-bias-gelu.wgsl.jinja",
47
  "bindings": ["x", { "arg": "bias", "elementType": "$scalar", "length": "$hidden" }, "y", "params"],
48
  "dispatch": {
 
61
  "passes": [
62
  {
63
  "id": "main",
64
+ "name": "FastGelu.Vec4",
65
  "shader": "elementwise-bias-gelu.wgsl.jinja",
66
  "bindings": ["x", "y", "params"],
67
  "dispatch": {
 
75
  {
76
  "id": "vec4_tail_bias",
77
  "priority": 20,
78
+ "when": ["biasOk"],
79
  "derive": {
80
  "vec4": false,
81
+ "vec4Tail": "numel(shapes.X) > 0",
82
  "hasBias": true,
83
  "hidden": "dim(shapes.X, ranks.X - 1) if dim(shapes.X, ranks.X - 1) > 0 else 1"
84
  },
85
  "passes": [
86
  {
87
  "id": "main",
88
+ "name": "FastGelu.Vec4TailBias",
89
  "shader": "elementwise-bias-gelu.wgsl.jinja",
90
+ "bindings": ["x_x", "bias", "y_y", "params_main"],
91
  "dispatch": {
92
  "x": "min(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
93
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
 
99
  {
100
  "id": "vec4_tail_no_bias",
101
  "priority": 15,
102
+ "when": ["noBiasOk"],
103
+ "derive": { "vec4": false, "vec4Tail": "numel(shapes.X) > 0", "hasBias": false },
104
  "passes": [
105
  {
106
  "id": "main",
107
+ "name": "FastGelu.Vec4Tail",
108
  "shader": "elementwise-bias-gelu.wgsl.jinja",
109
+ "bindings": ["x_x", "y_y", "params_main"],
110
  "dispatch": {
111
  "x": "min(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
112
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.X), 4)), (workgroupSize)), 65535)",
 
114
  }
115
  }
116
  ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
117
  }
118
  ]
119
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.FastGelu",
3
- "id": "_com_microsoft_fastgelu_webgpu_da3aabd",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,21 +8,19 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "uob2jmkWOUzFyVqQy0sJhB7JkQGvh+5ux8/7f+O6DYA=",
11
- "elementwise-bias-gelu.wgsl.jinja": "OW2nqbCYsSLMd8w5lKcNx4zjdkHi4i2kB+FFAKyHhFg=",
12
- "manifest.json": "nHE7yhM3Q5obzXWntT8sXryEtqYicMVDVSqdP+kwbLs=",
13
- "test.json": "oH8XR6itu84Ti3GZbKE8ImRXTl4N2njDh1qRw0+HhCs="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
  "variants": {
20
  "vec4_bias": ["elementwise-bias-gelu.wgsl.jinja"],
21
  "vec4_no_bias": ["elementwise-bias-gelu.wgsl.jinja"],
22
  "vec4_tail_bias": ["elementwise-bias-gelu.wgsl.jinja"],
23
- "vec4_tail_no_bias": ["elementwise-bias-gelu.wgsl.jinja"],
24
- "scalar_bias": ["elementwise-bias-gelu.wgsl.jinja"],
25
- "scalar_no_bias": ["elementwise-bias-gelu.wgsl.jinja"]
26
  }
27
  }
28
  }
 
1
  {
2
  "name": "com.microsoft.FastGelu",
3
+ "id": "_com_microsoft_fastgelu_webgpu_aa8ecda",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "uob2jmkWOUzFyVqQy0sJhB7JkQGvh+5ux8/7f+O6DYA=",
11
+ "elementwise-bias-gelu.wgsl.jinja": "RN55dM80C5gNCEsQC/n2ToecW/TW++sB2neIoraM05o=",
12
+ "manifest.json": "mUY8bymm2MAiLszkBwz6y1s+DT+ulKtK+MzMvKG8yyM=",
13
+ "test.json": "JgD1Mpetx475VrIbLzMW6x4yhiQIlCbYGKG1bTLptJ4="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
  "webgpu": {
18
+ "manifestSpec": "2.1",
19
  "variants": {
20
  "vec4_bias": ["elementwise-bias-gelu.wgsl.jinja"],
21
  "vec4_no_bias": ["elementwise-bias-gelu.wgsl.jinja"],
22
  "vec4_tail_bias": ["elementwise-bias-gelu.wgsl.jinja"],
23
+ "vec4_tail_no_bias": ["elementwise-bias-gelu.wgsl.jinja"]
 
 
24
  }
25
  }
26
  }
build/webgpu/test.json CHANGED
@@ -76,7 +76,7 @@
76
  {
77
  "name": "f32_scalar_no_bias_zero_sequence",
78
  "provenance": {
79
- "notes": "An empty bias-free input bypasses both vec4 routes and selects the scalar no-bias kernel."
80
  },
81
  "inputs": { "X": { "dtype": "float32", "shape": [1, 0, 4], "data": { "kind": "values", "values": [] } } },
82
  "outputs": { "Y": { "dtype": "float32", "shape": [1, 0, 4], "data": { "kind": "values", "values": [] } } }
@@ -442,6 +442,54 @@
442
  "data": { "kind": "values", "values": [-0.04540231, -0.15880801, 0.0, 0.84119199, 1.95459769] }
443
  }
444
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
445
  }
446
  ]
447
  }
 
76
  {
77
  "name": "f32_scalar_no_bias_zero_sequence",
78
  "provenance": {
79
+ "notes": "An empty float32 input with no bias must produce an empty output without reading an absent bias."
80
  },
81
  "inputs": { "X": { "dtype": "float32", "shape": [1, 0, 4], "data": { "kind": "values", "values": [] } } },
82
  "outputs": { "Y": { "dtype": "float32", "shape": [1, 0, 4], "data": { "kind": "values", "values": [] } } }
 
442
  "data": { "kind": "values", "values": [-0.04540231, -0.15880801, 0.0, 0.84119199, 1.95459769] }
443
  }
444
  }
445
+ },
446
+ {
447
+ "name": "tanh_near_zero_7",
448
+ "provenance": {
449
+ "notes": "Both sides of the tiny tanh input threshold, with scalar-tail and aligned vec4 shapes. Expected values use the mathematical tanh approximation."
450
+ },
451
+ "inputs": {
452
+ "X": {
453
+ "dtype": "float32",
454
+ "shape": [7],
455
+ "data": { "kind": "values", "values": [-0.000125, -0.00012, -1e-8, 0.0, 1e-8, 0.00012, 0.000125] }
456
+ }
457
+ },
458
+ "outputs": {
459
+ "Y": {
460
+ "dtype": "float32",
461
+ "shape": [7],
462
+ "data": {
463
+ "kind": "values",
464
+ "values": [-0.00006249376652688504, -0.000059994255231176074, -4.999999960105772e-9, 0.0, 5.000000039894228e-9, 0.00006000574476882393, 0.00006250623347311496]
465
+ },
466
+ "tolerance": 1e-10
467
+ }
468
+ }
469
+ },
470
+ {
471
+ "name": "tanh_near_zero_8",
472
+ "provenance": {
473
+ "notes": "Both sides of the tiny tanh input threshold, with scalar-tail and aligned vec4 shapes. Expected values use the mathematical tanh approximation."
474
+ },
475
+ "inputs": {
476
+ "X": {
477
+ "dtype": "float32",
478
+ "shape": [8],
479
+ "data": { "kind": "values", "values": [-0.000125, -0.00012, -1e-8, 0.0, 1e-8, 0.00012, 0.000125, 0.00013] }
480
+ }
481
+ },
482
+ "outputs": {
483
+ "Y": {
484
+ "dtype": "float32",
485
+ "shape": [8],
486
+ "data": {
487
+ "kind": "values",
488
+ "values": [-0.00006249376652688504, -0.000059994255231176074, -4.999999960105772e-9, 0.0, 5.000000039894228e-9, 0.00006000574476882393, 0.00006250623347311496, 0.0000650067421245197]
489
+ },
490
+ "tolerance": 1e-10
491
+ }
492
+ }
493
  }
494
  ]
495
  }