Xenova HF Staff commited on
Commit
f5120db
·
verified ·
1 Parent(s): 1092006

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -40,6 +40,12 @@ See the [ONNX Runtime `GemmaRotaryEmbedding` contrib-operator spec](https://gith
40
  | `T` | `float16` |
41
  | `U` | `float32` |
42
 
 
 
 
 
 
 
43
  ## Device requirements
44
 
45
  Every implementation variant requires `shader-f16`; the package has no variant-level fallback without that capability.
@@ -49,13 +55,13 @@ Every implementation variant requires `shader-f16`; the package has no variant-l
49
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
50
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
51
  - [`test.json`](build/webgpu/test.json) — correctness cases
52
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
53
  - [`gemma-rotary-embedding.wgsl.jinja`](build/webgpu/gemma-rotary-embedding.wgsl.jinja)
54
 
55
  ## Use with `@huggingface/kernels`
56
 
57
  ```sh
58
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
59
  ```
60
 
61
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
40
  | `T` | `float16` |
41
  | `U` | `float32` |
42
 
43
+ ## Implementation variants
44
+
45
+ One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
46
+
47
+ - `elementwise` — Elementwise rotation. Streams `vec4` words when each head's `(seq, dim)` plane holds a multiple of four elements, so no word straddles two heads; scalar elements otherwise.
48
+
49
  ## Device requirements
50
 
51
  Every implementation variant requires `shader-f16`; the package has no variant-level fallback without that capability.
 
55
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
56
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
57
  - [`test.json`](build/webgpu/test.json) — correctness cases
58
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
59
  - [`gemma-rotary-embedding.wgsl.jinja`](build/webgpu/gemma-rotary-embedding.wgsl.jinja)
60
 
61
  ## Use with `@huggingface/kernels`
62
 
63
  ```sh
64
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
65
  ```
66
 
67
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/gemma-rotary-embedding.wgsl.jinja CHANGED
@@ -1,4 +1,11 @@
1
- enable f16;
 
 
 
 
 
 
 
2
  {{ env.wgsl.resourceDeclarations }}
3
 
4
  // output1 = q * cos(emb) + q_rot * sin(emb); output2 applies the same
@@ -13,12 +20,7 @@ const ZERO: {{ scalar }} = {{ scalar }}(0.0);
13
 
14
  @compute @workgroup_size(WG, 1, 1)
15
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
16
- // 2D-folded flat index: gid.y carries the high bits past the
17
- // per-axis dispatch fold width. Reduces to gid.x when the dispatch does not fold.
18
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
19
- if (index >= params.count) {
20
- return;
21
- }
22
  // (batch, head, seq, dim) -> (batch, seq, dim): divide out the head-major stride and
23
  // keep the remainder, which is exactly the (seq, dim) offset shared by every head.
24
  let emb_index = (index / params.headSeqDim) * params.seqDim + index % params.seqDim;
 
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
  // output1 = q * cos(emb) + q_rot * sin(emb); output2 applies the same
 
20
 
21
  @compute @workgroup_size(WG, 1, 1)
22
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
23
+ {{ flat_index_2d("WG", "index") }}
 
 
 
 
 
24
  // (batch, head, seq, dim) -> (batch, seq, dim): divide out the head-major stride and
25
  // keep the remainder, which is exactly the (seq, dim) offset shared by every head.
26
  let emb_index = (index / params.headSeqDim) * params.seqDim + index % params.seqDim;
build/webgpu/manifest.json CHANGED
@@ -28,54 +28,15 @@
28
  "when": ["contract", "tunables.workgroupSize >= 1", "floor(tunables.workgroupSize) == tunables.workgroupSize", "tunables.workgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.workgroupSize <= device.limits.maxComputeWorkgroupSizeX"],
29
  "variants": [
30
  {
31
- "id": "vec4",
32
- "priority": 10,
33
- "when": ["vec4Ok"],
34
  "requires": { "features": ["shader-f16"] },
35
  "derive": {
36
- "vec4": true,
37
- "scalar": "dtypes.T",
38
- "vector": "\"vec4<f16>\"",
39
- "workgroupSize": "tunables.workgroupSize"
40
- },
41
- "passes": [
42
- {
43
- "id": "main",
44
- "name": "GemmaRotaryEmbedding.Vec4",
45
- "shader": "gemma-rotary-embedding.wgsl.jinja",
46
- "bindings": [
47
- { "arg": "embT", "name": "emb", "elementType": "vec4<f32>" },
48
- { "arg": "qT", "name": "q", "elementType": "$vector" },
49
- { "arg": "qRotT", "name": "q_rot", "elementType": "$vector" },
50
- { "arg": "kT", "name": "k", "elementType": "$vector" },
51
- { "arg": "kRotT", "name": "k_rot", "elementType": "$vector" },
52
- { "arg": "output1T", "name": "output1", "elementType": "$vector" },
53
- { "arg": "output2T", "name": "output2", "elementType": "$vector" },
54
- {
55
- "name": "params",
56
- "struct": [
57
- { "name": "count", "type": "u32", "value": "numel(shapes.qT) / 4" },
58
- { "name": "seqDim", "type": "u32", "value": "(seqLen * headDim) / 4" },
59
- { "name": "headSeqDim", "type": "u32", "value": "(numHeads * seqLen * headDim) / 4" }
60
- ]
61
- }
62
- ],
63
- "dispatch": {
64
- "x": "min(ceilDiv((numel(shapes.qT) / 4), (workgroupSize)), 65535)",
65
- "y": "ceilDiv(ceilDiv((numel(shapes.qT) / 4), (workgroupSize)), 65535)",
66
- "z": 1
67
- }
68
- }
69
- ]
70
- },
71
- {
72
- "id": "scalar",
73
- "priority": 0,
74
- "requires": { "features": ["shader-f16"] },
75
- "derive": {
76
- "vec4": false,
77
  "scalar": "dtypes.T",
78
  "vector": "\"vec4<f16>\"",
 
 
79
  "workgroupSize": "tunables.workgroupSize"
80
  },
81
  "passes": [
@@ -84,25 +45,25 @@
84
  "name": "GemmaRotaryEmbedding",
85
  "shader": "gemma-rotary-embedding.wgsl.jinja",
86
  "bindings": [
87
- { "arg": "embT", "name": "emb" },
88
- { "arg": "qT", "name": "q", "elementType": "$scalar" },
89
- { "arg": "qRotT", "name": "q_rot", "elementType": "$scalar" },
90
- { "arg": "kT", "name": "k", "elementType": "$scalar" },
91
- { "arg": "kRotT", "name": "k_rot", "elementType": "$scalar" },
92
- { "arg": "output1T", "name": "output1", "elementType": "$scalar" },
93
- { "arg": "output2T", "name": "output2", "elementType": "$scalar" },
94
  {
95
  "name": "params",
96
  "struct": [
97
- { "name": "count", "type": "u32", "value": "numel(shapes.qT)" },
98
- { "name": "seqDim", "type": "u32", "value": "(seqLen * headDim)" },
99
- { "name": "headSeqDim", "type": "u32", "value": "(numHeads * seqLen * headDim)" }
100
  ]
101
  }
102
  ],
103
  "dispatch": {
104
- "x": "min(ceilDiv((max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
105
- "y": "ceilDiv(ceilDiv((max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
106
  "z": 1
107
  }
108
  }
 
28
  "when": ["contract", "tunables.workgroupSize >= 1", "floor(tunables.workgroupSize) == tunables.workgroupSize", "tunables.workgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.workgroupSize <= device.limits.maxComputeWorkgroupSizeX"],
29
  "variants": [
30
  {
31
+ "id": "elementwise",
 
 
32
  "requires": { "features": ["shader-f16"] },
33
  "derive": {
34
+ "vec4": "vec4Ok",
35
+ "lanes": "4 if vec4Ok else 1",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36
  "scalar": "dtypes.T",
37
  "vector": "\"vec4<f16>\"",
38
+ "element": "vector if vec4Ok else scalar",
39
+ "embElement": "\"vec4<f32>\" if vec4Ok else \"f32\"",
40
  "workgroupSize": "tunables.workgroupSize"
41
  },
42
  "passes": [
 
45
  "name": "GemmaRotaryEmbedding",
46
  "shader": "gemma-rotary-embedding.wgsl.jinja",
47
  "bindings": [
48
+ { "arg": "embT", "name": "emb", "elementType": "$embElement" },
49
+ { "arg": "qT", "name": "q", "elementType": "$element" },
50
+ { "arg": "qRotT", "name": "q_rot", "elementType": "$element" },
51
+ { "arg": "kT", "name": "k", "elementType": "$element" },
52
+ { "arg": "kRotT", "name": "k_rot", "elementType": "$element" },
53
+ { "arg": "output1T", "name": "output1", "elementType": "$element" },
54
+ { "arg": "output2T", "name": "output2", "elementType": "$element" },
55
  {
56
  "name": "params",
57
  "struct": [
58
+ { "name": "count", "type": "u32", "value": "numel(shapes.qT) / lanes" },
59
+ { "name": "seqDim", "type": "u32", "value": "(seqLen * headDim) / lanes" },
60
+ { "name": "headSeqDim", "type": "u32", "value": "(numHeads * seqLen * headDim) / lanes" }
61
  ]
62
  }
63
  ],
64
  "dispatch": {
65
+ "x": "min(ceilDiv((numel(shapes.qT) / 4 if vec4Ok else max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
66
+ "y": "ceilDiv(ceilDiv((numel(shapes.qT) / 4 if vec4Ok else max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
67
  "z": 1
68
  }
69
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.GemmaRotaryEmbedding",
3
- "id": "_com_microsoft_gemmarotaryembedding_webgpu_e5a6aca",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,14 +8,11 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "E1026CMMv9Pppl2+JpGoj3dWCSmk8tZUEwBYo3UCiD4=",
11
- "gemma-rotary-embedding.wgsl.jinja": "t4wKCJ/E3156iIqQTq1d2XGqMBA83kI0EAADYck7NI0=",
12
- "manifest.json": "41x/WJ1/AyRVHsYgM3nIqECf6Z8zix9MBzVgLIaLILE=",
13
- "test.json": "WwfG+ffVK8tAHKvZCOkixu97BOEIQjEdR4NSOSLA9dc="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
- "webgpu": {
18
- "manifestSpec": "2.0",
19
- "variants": { "vec4": ["gemma-rotary-embedding.wgsl.jinja"], "scalar": ["gemma-rotary-embedding.wgsl.jinja"] }
20
- }
21
  }
 
1
  {
2
  "name": "com.microsoft.GemmaRotaryEmbedding",
3
+ "id": "_com_microsoft_gemmarotaryembedding_webgpu_bb6bdf4",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "E1026CMMv9Pppl2+JpGoj3dWCSmk8tZUEwBYo3UCiD4=",
11
+ "gemma-rotary-embedding.wgsl.jinja": "J7OaLWdNRVSQJN7334F9gTgV5OrdWUywfjXx5rPnHr4=",
12
+ "manifest.json": "QJJdYbfJe3E3SRzXSXLz3IuLFfsthXglsTWLtCQWh4U=",
13
+ "test.json": "0qR3IsgdZC/lwA1WbH1QDFYTijEjsyo3t++UOJnt6VY="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
+ "webgpu": { "manifestSpec": "2.1", "variants": { "elementwise": ["gemma-rotary-embedding.wgsl.jinja"] } }
 
 
 
18
  }
build/webgpu/test.json CHANGED
@@ -201,7 +201,7 @@
201
  {
202
  "name": "f16_vec4_dim6_seq4_straddles_heads",
203
  "provenance": {
204
- "notes": "headDim 6 is not a multiple of four but seqLen * headDim = 24 is: the vec4 route admits it and its vec4s straddle head-dim boundaries within one (batch, seq) block."
205
  },
206
  "inputs": {
207
  "embT": {
 
201
  {
202
  "name": "f16_vec4_dim6_seq4_straddles_heads",
203
  "provenance": {
204
+ "notes": "Head dim 6 is not a multiple of four, but seq_len x head_dim = 4x6 = 24 is; checks four-element grouping that straddles head-dim boundaries within one (batch, seq) block rather than aligning to them."
205
  },
206
  "inputs": {
207
  "embT": {