Xenova HF Staff commited on
Commit
77f8a08
·
verified ·
1 Parent(s): 7d7921b

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -1,3 +1,84 @@
1
  ---
 
2
  license: apache-2.0
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ library_name: kernels
3
  license: apache-2.0
4
+ tags:
5
+ - kernel
6
+ - webgpu
7
+ - wgsl
8
  ---
9
+ # com.microsoft.RotaryEmbedding
10
+
11
+ `com.microsoft` · ONNX Runtime contrib operator · contrib since_version 1
12
+
13
+ ## Description
14
+
15
+ Rotary positional embedding (RoPE): each head's embedding vector is rotated with the `cos_cache` and `sin_cache` rows selected by `position_ids`, which is either a single base offset (token `s` reads row `position_ids[0] + s`) or a `(batch_size, sequence_length)` table. `input` is rank 3 `(batch_size, sequence_length, hidden_size)` or rank 4 `(batch_size, num_heads, sequence_length, head_size)`. `rotary_embedding_dim` rotates a prefix and copies the tail unchanged. Rotation arithmetic is float32 with one narrowing store. Bfloat16 and non-default `scale` are not implemented.
16
+
17
+ See the [ONNX Runtime `RotaryEmbedding` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.RotaryEmbedding) for the reference semantics.
18
+
19
+ ## Inputs
20
+
21
+ | Name | Upstream name | Logical dtype | WebGPU storage | Rank | Shape | Description | Presence |
22
+ | --- | --- | --- | --- | --- | --- | --- | --- |
23
+ | `x` | `input` | `T` | same as logical dtype | — | — | Input token embeddings, shaped `(batch_size, sequence_length, hidden_size)` at rank 3 or `(batch_size, num_heads, sequence_length, head_size)` at rank 4. At rank 3 the head size comes from `num_heads` when that attribute is positive and from `2 * cos_cache.shape[1]` otherwise; at rank 4 both the head count and head size are read from the shape. | required |
24
+ | `positionIds` | `position_ids` | `M` | `uint32` | — | — | Logical int64 cache-row selector in either upstream format: a scalar or one-element vector holding a base offset, so token `s` reads row `position_ids[0] + s`; or a `(batch_size, sequence_length)` table read per token. Valid positions are non-negative rows of the caches and use uint32 WebGPU storage, so a negative value is rejected at the host boundary. | required |
25
+ | `cos` | `cos_cache` | `T` | same as logical dtype | `2` | — | Precomputed cosine values of shape `(max_sequence_length, rotary_dim / 2)`, where `rotary_dim` is `rotary_embedding_dim` when that is positive and the head size otherwise. | required |
26
+ | `sin` | `sin_cache` | `T` | same as logical dtype | `2` | — | Precomputed sine values with the same shape and type as `cos_cache`. | required |
27
+
28
+ ## Outputs
29
+
30
+ | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
31
+ | --- | --- | --- | --- | --- | --- | --- |
32
+ | `y` | `output` | `T` | same as `x` | same as `x` | Rotary-position-encoded tensor with the same shape and type as `input`. | required |
33
+
34
+ ## Attributes
35
+
36
+ Default values (overridable per request):
37
+
38
+ | Attribute | Default | Description |
39
+ | --- | --- | --- |
40
+ | `interleaved` | `0` | Set to 1 to pair adjacent even/odd elements, or 0 to pair each element of the first half of the rotary window with the matching element of the second half. Default is 0. |
41
+ | `is_packed_batching` | `0` | Ragged (packed) batch inputs. Its only upstream effect is to lift the `sequence_length <= max_sequence_length` bound, and this implementation never imposes that bound: every gathered row is required to index the caches whatever the sequence length. Default is 0. |
42
+ | `num_heads` | `0` | Number of attention heads. Default is 0, which asks the rank-3 path to take the head size from the cache width instead; a positive value is required whenever `rotary_embedding_dim` is nonzero. At rank 4 the head count comes from the input shape and this attribute is not consulted. |
43
+ | `rotary_embedding_dim` | `0` | Positive even count of leading head-dimension elements to rotate; the remaining tail is copied unchanged. Default is 0, meaning the whole head dimension, which must then be even. An odd value is rejected. |
44
+ | `scale` | `1` | Declared scale for the gathered rotation. No ONNX Runtime provider applies it, so the default and only accepted value is 1.0 and a caller's other value is rejected rather than silently discarded. |
45
+
46
+ ## Type constraints
47
+
48
+ | Variable | Allowed dtypes |
49
+ | --- | --- |
50
+ | `T` | `float32`, `float16` |
51
+ | `M` | `int64` |
52
+
53
+ ## Files
54
+
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
+ - [`rotary-embedding-slices.wgsl.jinja`](build/webgpu/rotary-embedding-slices.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.
68
+
69
+ The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
70
+ It follows the `v1` branch as fixes land. To pin exact artifact bytes, pass a 40-character commit `revision` instead of `version`.
71
+
72
+ Replace each `*Data` placeholder with a typed array containing the corresponding input data.
73
+
74
+ ```js
75
+ import { getKernel } from "@huggingface/kernels";
76
+
77
+ const kernel = await getKernel("webgpu-kernels/com.microsoft.RotaryEmbedding", { version: 1 });
78
+ const { y } = await kernel({
79
+ x: { data: xData, shape: [1, 2, 18] },
80
+ positionIds: { data: positionIdsData, shape: [1, 2] },
81
+ cos: { data: cosData, shape: [4, 3] },
82
+ sin: { data: sinData, shape: [4, 3] },
83
+ });
84
+ ```
build/webgpu/bench.json ADDED
@@ -0,0 +1,442 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "tunableSpace": { "WORKGROUP_SIZE": [64, 128, 256, 512], "ITEMS_PER_LANE": [1, 2, 4] },
3
+ "cases": [
4
+ {
5
+ "name": "llama3-8b-prefill-s512-q-h32-d128-f32",
6
+ "preset": "smoke",
7
+ "attrs": { "interleaved": 0 },
8
+ "provenance": {
9
+ "notes": "512-token prefill of the query projection: rank-3 packed heads with a rank-2 cache and a per-token position table. Read-once, write-once traffic.",
10
+ "model": "meta-llama/Meta-Llama-3-8B"
11
+ },
12
+ "vars": { "dtype": "float32", "elems": 2097152, "cacheElems": 32768 },
13
+ "inputs": {
14
+ "x": { "shape": [1, 512, 4096], "dtype": "float32", "dist": "normal", "seed": 4010, "scale": 1 },
15
+ "positionIds": {
16
+ "shape": [1, 512],
17
+ "dtype": "uint32",
18
+ "dist": "linearMod",
19
+ "seed": 4010,
20
+ "step": 1,
21
+ "offset": 0,
22
+ "mod": 4096
23
+ },
24
+ "cos": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4011, "scale": 0.1, "offset": 0.9 },
25
+ "sin": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4012, "scale": 0.1 }
26
+ },
27
+ "outputs": { "y": { "shape": [1, 512, 4096], "dtype": "float32" } },
28
+ "bench": {
29
+ "primary": true,
30
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
31
+ }
32
+ },
33
+ {
34
+ "name": "llama3-8b-prefill-s512-q-h32-d128-f16",
35
+ "preset": "smoke",
36
+ "attrs": { "interleaved": 0 },
37
+ "provenance": {
38
+ "notes": "The float16 prefill of the same projection, the dtype shipped model graphs use.",
39
+ "model": "meta-llama/Meta-Llama-3-8B"
40
+ },
41
+ "vars": { "dtype": "float16", "elems": 2097152, "cacheElems": 32768 },
42
+ "inputs": {
43
+ "x": { "shape": [1, 512, 4096], "dtype": "float16", "dist": "normal", "seed": 4020, "scale": 1 },
44
+ "positionIds": {
45
+ "shape": [1, 512],
46
+ "dtype": "uint32",
47
+ "dist": "linearMod",
48
+ "seed": 4020,
49
+ "step": 1,
50
+ "offset": 0,
51
+ "mod": 4096
52
+ },
53
+ "cos": { "shape": [4096, 64], "dtype": "float16", "dist": "uniform", "seed": 4021, "scale": 0.1, "offset": 0.9 },
54
+ "sin": { "shape": [4096, 64], "dtype": "float16", "dist": "uniform", "seed": 4022, "scale": 0.1 }
55
+ },
56
+ "outputs": { "y": { "shape": [1, 512, 4096], "dtype": "float16" } },
57
+ "bench": {
58
+ "primary": true,
59
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
60
+ }
61
+ },
62
+ {
63
+ "name": "llama3-8b-prefill-s512-q-h32-d128-interleaved-f32",
64
+ "preset": "smoke",
65
+ "attrs": { "interleaved": 1 },
66
+ "provenance": {
67
+ "notes": "The interleaved pairing at the same shape. One of the two ONNX Runtime fusion paths emits interleaved=1, so both layouts occur in real exports.",
68
+ "model": "meta-llama/Meta-Llama-3-8B"
69
+ },
70
+ "vars": { "dtype": "float32", "elems": 2097152, "cacheElems": 32768 },
71
+ "inputs": {
72
+ "x": { "shape": [1, 512, 4096], "dtype": "float32", "dist": "normal", "seed": 4030, "scale": 1 },
73
+ "positionIds": {
74
+ "shape": [1, 512],
75
+ "dtype": "uint32",
76
+ "dist": "linearMod",
77
+ "seed": 4030,
78
+ "step": 1,
79
+ "offset": 0,
80
+ "mod": 4096
81
+ },
82
+ "cos": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4031, "scale": 0.1, "offset": 0.9 },
83
+ "sin": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4032, "scale": 0.1 }
84
+ },
85
+ "outputs": { "y": { "shape": [1, 512, 4096], "dtype": "float32" } },
86
+ "bench": {
87
+ "primary": false,
88
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
89
+ }
90
+ },
91
+ {
92
+ "name": "llama3-8b-decode-q-h32-d128-f32",
93
+ "preset": "smoke",
94
+ "attrs": { "interleaved": 0 },
95
+ "provenance": {
96
+ "notes": "One decoded token through a format-0 base offset, the shape a decode loop emits. 32 KB of traffic, so this measures launch latency.",
97
+ "model": "meta-llama/Meta-Llama-3-8B"
98
+ },
99
+ "vars": { "dtype": "float32", "elems": 4096, "cacheElems": 64 },
100
+ "inputs": {
101
+ "x": { "shape": [1, 1, 4096], "dtype": "float32", "dist": "normal", "seed": 4040, "scale": 1 },
102
+ "positionIds": { "shape": [1], "dtype": "uint32", "data": { "kind": "values", "values": [1024] } },
103
+ "cos": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4041, "scale": 0.1, "offset": 0.9 },
104
+ "sin": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4042, "scale": 0.1 }
105
+ },
106
+ "outputs": { "y": { "shape": [1, 1, 4096], "dtype": "float32" } },
107
+ "bench": {
108
+ "primary": false,
109
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
110
+ }
111
+ },
112
+ {
113
+ "name": "llama3-8b-decode-k-h8-d128-f32",
114
+ "preset": "smoke",
115
+ "attrs": { "interleaved": 0 },
116
+ "provenance": {
117
+ "notes": "The grouped-query key projection of the same decode step: eight heads, the smallest dispatch this operator sees per token.",
118
+ "model": "meta-llama/Meta-Llama-3-8B"
119
+ },
120
+ "vars": { "dtype": "float32", "elems": 1024, "cacheElems": 64 },
121
+ "inputs": {
122
+ "x": { "shape": [1, 1, 1024], "dtype": "float32", "dist": "normal", "seed": 4050, "scale": 1 },
123
+ "positionIds": { "shape": [1], "dtype": "uint32", "data": { "kind": "values", "values": [1024] } },
124
+ "cos": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4051, "scale": 0.1, "offset": 0.9 },
125
+ "sin": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4052, "scale": 0.1 }
126
+ },
127
+ "outputs": { "y": { "shape": [1, 1, 1024], "dtype": "float32" } },
128
+ "bench": {
129
+ "primary": false,
130
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
131
+ }
132
+ },
133
+ {
134
+ "name": "qwen3-4b-prefill-s512-rank4-h32-d128-f16",
135
+ "preset": "smoke",
136
+ "attrs": { "interleaved": 0 },
137
+ "provenance": {
138
+ "notes": "Rank-4 BNSH prefill, the transposed layout the ONNX Runtime attention fusion emits after a Transpose.",
139
+ "model": "Qwen/Qwen3-4B"
140
+ },
141
+ "vars": { "dtype": "float16", "elems": 2097152, "cacheElems": 32768 },
142
+ "inputs": {
143
+ "x": { "shape": [1, 32, 512, 128], "dtype": "float16", "dist": "normal", "seed": 4060, "scale": 1 },
144
+ "positionIds": {
145
+ "shape": [1, 512],
146
+ "dtype": "uint32",
147
+ "dist": "linearMod",
148
+ "seed": 4060,
149
+ "step": 1,
150
+ "offset": 0,
151
+ "mod": 4096
152
+ },
153
+ "cos": { "shape": [4096, 64], "dtype": "float16", "dist": "uniform", "seed": 4061, "scale": 0.1, "offset": 0.9 },
154
+ "sin": { "shape": [4096, 64], "dtype": "float16", "dist": "uniform", "seed": 4062, "scale": 0.1 }
155
+ },
156
+ "outputs": { "y": { "shape": [1, 32, 512, 128], "dtype": "float16" } },
157
+ "bench": {
158
+ "primary": false,
159
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
160
+ }
161
+ },
162
+ {
163
+ "name": "phi2-prefill-s512-rotary32-d80-f32",
164
+ "preset": "smoke",
165
+ "attrs": { "interleaved": 0, "rotary_embedding_dim": 32, "num_heads": 32 },
166
+ "provenance": {
167
+ "notes": "Partial rotary: 32 of 80 head elements rotate, so 60 percent of the traffic is an untouched tail the same invocations copy at full vector width.",
168
+ "model": "microsoft/phi-2"
169
+ },
170
+ "vars": { "dtype": "float32", "elems": 1310720, "cacheElems": 8192 },
171
+ "inputs": {
172
+ "x": { "shape": [1, 512, 2560], "dtype": "float32", "dist": "normal", "seed": 4070, "scale": 1 },
173
+ "positionIds": {
174
+ "shape": [1, 512],
175
+ "dtype": "uint32",
176
+ "dist": "linearMod",
177
+ "seed": 4070,
178
+ "step": 1,
179
+ "offset": 0,
180
+ "mod": 4096
181
+ },
182
+ "cos": { "shape": [4096, 16], "dtype": "float32", "dist": "uniform", "seed": 4071, "scale": 0.1, "offset": 0.9 },
183
+ "sin": { "shape": [4096, 16], "dtype": "float32", "dist": "uniform", "seed": 4072, "scale": 0.1 }
184
+ },
185
+ "outputs": { "y": { "shape": [1, 512, 2560], "dtype": "float32" } },
186
+ "bench": {
187
+ "primary": false,
188
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
189
+ }
190
+ },
191
+ {
192
+ "name": "phi3-mini-prefill-s512-d96-h32-f16",
193
+ "preset": "smoke",
194
+ "attrs": { "interleaved": 0 },
195
+ "provenance": {
196
+ "notes": "Head size 96 with a 48-wide cache: a non-power-of-two head that still clears the four-pair route's gates.",
197
+ "model": "microsoft/Phi-3-mini-4k-instruct"
198
+ },
199
+ "vars": { "dtype": "float16", "elems": 1572864, "cacheElems": 24576 },
200
+ "inputs": {
201
+ "x": { "shape": [1, 512, 3072], "dtype": "float16", "dist": "normal", "seed": 4080, "scale": 1 },
202
+ "positionIds": {
203
+ "shape": [1, 512],
204
+ "dtype": "uint32",
205
+ "dist": "linearMod",
206
+ "seed": 4080,
207
+ "step": 1,
208
+ "offset": 0,
209
+ "mod": 4096
210
+ },
211
+ "cos": { "shape": [4096, 48], "dtype": "float16", "dist": "uniform", "seed": 4081, "scale": 0.1, "offset": 0.9 },
212
+ "sin": { "shape": [4096, 48], "dtype": "float16", "dist": "uniform", "seed": 4082, "scale": 0.1 }
213
+ },
214
+ "outputs": { "y": { "shape": [1, 512, 3072], "dtype": "float16" } },
215
+ "bench": {
216
+ "primary": false,
217
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
218
+ }
219
+ },
220
+ {
221
+ "name": "pair-route-prefill-s512-rotary126-d128-h32-f32",
222
+ "preset": "smoke",
223
+ "attrs": { "interleaved": 0, "rotary_embedding_dim": 126, "num_heads": 32 },
224
+ "provenance": {
225
+ "notes": "A rotary window of 126 gives 63 rotation pairs and leaves two unrotated elements in each 128-wide head, at the same geometry as the paired Llama prefill case."
226
+ },
227
+ "vars": { "dtype": "float32", "elems": 2097152, "cacheElems": 32256 },
228
+ "inputs": {
229
+ "x": { "shape": [1, 512, 4096], "dtype": "float32", "dist": "normal", "seed": 4090, "scale": 1 },
230
+ "positionIds": {
231
+ "shape": [1, 512],
232
+ "dtype": "uint32",
233
+ "dist": "linearMod",
234
+ "seed": 4090,
235
+ "step": 1,
236
+ "offset": 0,
237
+ "mod": 4096
238
+ },
239
+ "cos": { "shape": [4096, 63], "dtype": "float32", "dist": "uniform", "seed": 4091, "scale": 0.1, "offset": 0.9 },
240
+ "sin": { "shape": [4096, 63], "dtype": "float32", "dist": "uniform", "seed": 4092, "scale": 0.1 }
241
+ },
242
+ "outputs": { "y": { "shape": [1, 512, 4096], "dtype": "float32" } },
243
+ "bench": {
244
+ "primary": false,
245
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
246
+ }
247
+ },
248
+ {
249
+ "name": "llama3-8b-prefill-s1536-q-h32-d128-f32",
250
+ "preset": "model",
251
+ "attrs": { "interleaved": 0 },
252
+ "provenance": {
253
+ "notes": "A 1536-token prefill, sized so its working set matches the elementwise control of the same working set.",
254
+ "model": "meta-llama/Meta-Llama-3-8B"
255
+ },
256
+ "vars": { "dtype": "float32", "elems": 6291456, "cacheElems": 98304 },
257
+ "inputs": {
258
+ "x": { "shape": [1, 1536, 4096], "dtype": "float32", "dist": "normal", "seed": 4100, "scale": 1 },
259
+ "positionIds": {
260
+ "shape": [1, 1536],
261
+ "dtype": "uint32",
262
+ "dist": "linearMod",
263
+ "seed": 4100,
264
+ "step": 1,
265
+ "offset": 0,
266
+ "mod": 4096
267
+ },
268
+ "cos": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4101, "scale": 0.1, "offset": 0.9 },
269
+ "sin": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4102, "scale": 0.1 }
270
+ },
271
+ "outputs": { "y": { "shape": [1, 1536, 4096], "dtype": "float32" } },
272
+ "bench": {
273
+ "primary": false,
274
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
275
+ }
276
+ },
277
+ {
278
+ "name": "pair-route-prefill-s2048-rotary126-d128-h32-f32",
279
+ "preset": "model",
280
+ "attrs": { "interleaved": 0, "rotary_embedding_dim": 126, "num_heads": 32 },
281
+ "provenance": {
282
+ "notes": "The single-pair route at the same 67 MB working set as the four-pair route's model case, so the two schedules are directly comparable."
283
+ },
284
+ "vars": { "dtype": "float32", "elems": 8388608, "cacheElems": 129024 },
285
+ "inputs": {
286
+ "x": { "shape": [1, 2048, 4096], "dtype": "float32", "dist": "normal", "seed": 4110, "scale": 1 },
287
+ "positionIds": {
288
+ "shape": [1, 2048],
289
+ "dtype": "uint32",
290
+ "dist": "linearMod",
291
+ "seed": 4110,
292
+ "step": 1,
293
+ "offset": 0,
294
+ "mod": 8192
295
+ },
296
+ "cos": { "shape": [8192, 63], "dtype": "float32", "dist": "uniform", "seed": 4111, "scale": 0.1, "offset": 0.9 },
297
+ "sin": { "shape": [8192, 63], "dtype": "float32", "dist": "uniform", "seed": 4112, "scale": 0.1 }
298
+ },
299
+ "outputs": { "y": { "shape": [1, 2048, 4096], "dtype": "float32" } },
300
+ "bench": {
301
+ "primary": false,
302
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
303
+ }
304
+ },
305
+ {
306
+ "name": "qwen3-4b-decode-rank4-h32-d128-f32",
307
+ "preset": "smoke",
308
+ "attrs": { "interleaved": 0 },
309
+ "provenance": {
310
+ "notes": "One decoded token in the rank-4 layout through a format-0 base offset: 32 head blocks of one token each, the smallest rank-4 dispatch.",
311
+ "model": "Qwen/Qwen3-4B"
312
+ },
313
+ "vars": { "dtype": "float32", "elems": 4096, "cacheElems": 64 },
314
+ "inputs": {
315
+ "x": { "shape": [1, 32, 1, 128], "dtype": "float32", "dist": "normal", "seed": 4120, "scale": 1 },
316
+ "positionIds": { "shape": [1], "dtype": "uint32", "data": { "kind": "values", "values": [1024] } },
317
+ "cos": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4121, "scale": 0.1, "offset": 0.9 },
318
+ "sin": { "shape": [4096, 64], "dtype": "float32", "dist": "uniform", "seed": 4122, "scale": 0.1 }
319
+ },
320
+ "outputs": { "y": { "shape": [1, 32, 1, 128], "dtype": "float32" } },
321
+ "bench": {
322
+ "primary": false,
323
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
324
+ }
325
+ },
326
+ {
327
+ "name": "pair-route-rank4-prefill-s512-rotary126-f32",
328
+ "preset": "smoke",
329
+ "attrs": { "interleaved": 0, "rotary_embedding_dim": 126, "num_heads": 32 },
330
+ "provenance": {
331
+ "notes": "The single-pair route in the rank-4 layout, where the head block shares one gathered cache row across its heads exactly as the vector route does."
332
+ },
333
+ "vars": { "dtype": "float32", "elems": 2097152, "cacheElems": 32256 },
334
+ "inputs": {
335
+ "x": { "shape": [1, 32, 512, 128], "dtype": "float32", "dist": "normal", "seed": 4130, "scale": 1 },
336
+ "positionIds": {
337
+ "shape": [1, 512],
338
+ "dtype": "uint32",
339
+ "dist": "linearMod",
340
+ "seed": 4130,
341
+ "step": 1,
342
+ "offset": 0,
343
+ "mod": 4096
344
+ },
345
+ "cos": { "shape": [4096, 63], "dtype": "float32", "dist": "uniform", "seed": 4131, "scale": 0.1, "offset": 0.9 },
346
+ "sin": { "shape": [4096, 63], "dtype": "float32", "dist": "uniform", "seed": 4132, "scale": 0.1 }
347
+ },
348
+ "outputs": { "y": { "shape": [1, 32, 512, 128], "dtype": "float32" } },
349
+ "bench": {
350
+ "primary": false,
351
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
352
+ }
353
+ },
354
+ {
355
+ "name": "llama3-8b-prefill-s4096-q-h32-d128-f16",
356
+ "preset": "model",
357
+ "attrs": { "interleaved": 0 },
358
+ "provenance": {
359
+ "notes": "A 4096-token prefill: a 67 MB working set, far past cache, so the bandwidth figure is a DRAM figure.",
360
+ "model": "meta-llama/Meta-Llama-3-8B"
361
+ },
362
+ "vars": { "dtype": "float16", "elems": 16777216, "cacheElems": 262144 },
363
+ "inputs": {
364
+ "x": { "shape": [1, 4096, 4096], "dtype": "float16", "dist": "normal", "seed": 4140, "scale": 1 },
365
+ "positionIds": {
366
+ "shape": [1, 4096],
367
+ "dtype": "uint32",
368
+ "dist": "linearMod",
369
+ "seed": 4140,
370
+ "step": 1,
371
+ "offset": 0,
372
+ "mod": 8192
373
+ },
374
+ "cos": { "shape": [8192, 64], "dtype": "float16", "dist": "uniform", "seed": 4141, "scale": 0.1, "offset": 0.9 },
375
+ "sin": { "shape": [8192, 64], "dtype": "float16", "dist": "uniform", "seed": 4142, "scale": 0.1 }
376
+ },
377
+ "outputs": { "y": { "shape": [1, 4096, 4096], "dtype": "float16" } },
378
+ "bench": {
379
+ "primary": false,
380
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
381
+ }
382
+ },
383
+ {
384
+ "name": "llama3-8b-prefill-s2048-q-h32-d128-f32",
385
+ "preset": "model",
386
+ "attrs": { "interleaved": 0 },
387
+ "provenance": {
388
+ "notes": "The float32 counterpart at the same 67 MB working set.",
389
+ "model": "meta-llama/Meta-Llama-3-8B"
390
+ },
391
+ "vars": { "dtype": "float32", "elems": 8388608, "cacheElems": 131072 },
392
+ "inputs": {
393
+ "x": { "shape": [1, 2048, 4096], "dtype": "float32", "dist": "normal", "seed": 4150, "scale": 1 },
394
+ "positionIds": {
395
+ "shape": [1, 2048],
396
+ "dtype": "uint32",
397
+ "dist": "linearMod",
398
+ "seed": 4150,
399
+ "step": 1,
400
+ "offset": 0,
401
+ "mod": 8192
402
+ },
403
+ "cos": { "shape": [8192, 64], "dtype": "float32", "dist": "uniform", "seed": 4151, "scale": 0.1, "offset": 0.9 },
404
+ "sin": { "shape": [8192, 64], "dtype": "float32", "dist": "uniform", "seed": 4152, "scale": 0.1 }
405
+ },
406
+ "outputs": { "y": { "shape": [1, 2048, 4096], "dtype": "float32" } },
407
+ "bench": {
408
+ "primary": false,
409
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
410
+ }
411
+ },
412
+ {
413
+ "name": "qwen3-4b-prefill-s2048-rank4-h32-d128-f32",
414
+ "preset": "model",
415
+ "attrs": { "interleaved": 0 },
416
+ "provenance": {
417
+ "notes": "The rank-4 layout at the same 67 MB working set, where the pinned (batch, head) slice carries 2048 tokens.",
418
+ "model": "Qwen/Qwen3-4B"
419
+ },
420
+ "vars": { "dtype": "float32", "elems": 8388608, "cacheElems": 131072 },
421
+ "inputs": {
422
+ "x": { "shape": [1, 32, 2048, 128], "dtype": "float32", "dist": "normal", "seed": 4160, "scale": 1 },
423
+ "positionIds": {
424
+ "shape": [1, 2048],
425
+ "dtype": "uint32",
426
+ "dist": "linearMod",
427
+ "seed": 4160,
428
+ "step": 1,
429
+ "offset": 0,
430
+ "mod": 8192
431
+ },
432
+ "cos": { "shape": [8192, 64], "dtype": "float32", "dist": "uniform", "seed": 4161, "scale": 0.1, "offset": 0.9 },
433
+ "sin": { "shape": [8192, 64], "dtype": "float32", "dist": "uniform", "seed": 4162, "scale": 0.1 }
434
+ },
435
+ "outputs": { "y": { "shape": [1, 32, 2048, 128], "dtype": "float32" } },
436
+ "bench": {
437
+ "primary": false,
438
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.cacheElems * 2) * dtypeBytes(args.dtype)" }]
439
+ }
440
+ }
441
+ ]
442
+ }
build/webgpu/manifest.json ADDED
@@ -0,0 +1,146 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "domain": "com.microsoft",
3
+ "name": "RotaryEmbedding",
4
+ "sinceVersion": 1,
5
+ "inputs": {
6
+ "x": { "onnx": "input", "dtype": "T" },
7
+ "positionIds": { "onnx": "position_ids", "dtype": "M", "storage": "uint32", "narrowing": "checked" },
8
+ "cos": { "onnx": "cos_cache", "dtype": "T", "rank": 2 },
9
+ "sin": { "onnx": "sin_cache", "dtype": "T", "rank": 2 }
10
+ },
11
+ "outputs": { "y": { "onnx": "output", "dtype": "T", "rank": "ranks.x", "shape": "shapes.x" } },
12
+ "attributes": {
13
+ "interleaved": { "default": 0 },
14
+ "is_packed_batching": { "default": 0 },
15
+ "num_heads": { "default": 0 },
16
+ "rotary_embedding_dim": { "default": 0 },
17
+ "scale": { "default": 1 }
18
+ },
19
+ "attributeConstraints": {
20
+ "interleaved": { "values": [0, 1] },
21
+ "is_packed_batching": { "values": [0, 1] },
22
+ "scale": { "values": [1] }
23
+ },
24
+ "typeConstraints": { "T": ["float32", "float16"], "M": ["int64"] },
25
+ "tunables": { "WORKGROUP_SIZE": { "default": 256 }, "ITEMS_PER_LANE": { "default": 2 } },
26
+ "derive": {
27
+ "foldedDispatchCapacity": "min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
28
+ "rank3": "ranks.x == 3",
29
+ "rank4": "ranks.x == 4",
30
+ "rankOk": "rank3 or rank4",
31
+ "cacheShapeOk": "ranks.cos == 2 and ranks.sin == 2 and sameShape(shapes.cos, shapes.sin)",
32
+ "cacheWidth": "dim(shapes.cos, 1) if cacheShapeOk else 0",
33
+ "batchSize": "dim(shapes.x, 0) if rankOk else 0",
34
+ "seqLength": "(dim(shapes.x, 1) if rank3 else dim(shapes.x, 2)) if rankOk else 0",
35
+ "hiddenSize": "dim(shapes.x, 2) if rank3 else 0",
36
+ "rank3HeadSize": "(hiddenSize / attrs.num_heads if attrs.num_heads > 0 and hiddenSize % max(1, attrs.num_heads) == 0 else 2 * cacheWidth) if rank3 else 0",
37
+ "headSize": "rank3HeadSize if rank3 else (dim(shapes.x, 3) if rank4 else 0)",
38
+ "numHeads": "(hiddenSize / headSize if headSize > 0 and hiddenSize % max(1, headSize) == 0 else 0) if rank3 else (dim(shapes.x, 1) if rank4 else 0)",
39
+ "rotaryDim": "attrs.rotary_embedding_dim if attrs.rotary_embedding_dim > 0 else headSize",
40
+ "halfRotaryDim": "rotaryDim / 2 if rotaryDim % 2 == 0 else 0",
41
+ "geometryOk": "rankOk and cacheShapeOk and headSize > 0 and numHeads > 0 and batchSize > 0 and seqLength > 0 and rotaryDim > 0 and rotaryDim % 2 == 0 and rotaryDim <= headSize and cacheWidth == halfRotaryDim",
42
+ "attrsOk": "attrs.num_heads >= 0 and attrs.rotary_embedding_dim >= 0 and (attrs.rotary_embedding_dim == 0 or attrs.num_heads > 0)",
43
+ "contractOk": "f16Ok(dtypes.T) and geometryOk and attrsOk and sameShape(shapes.x, shapes.y)",
44
+ "itemsPerLane": "tunables.ITEMS_PER_LANE if (rank3 or numHeads % max(1, tunables.ITEMS_PER_LANE) == 0) else 1",
45
+ "headBlocks": "ceilDiv(numHeads, max(1, itemsPerLane)) if rank4 else 1",
46
+ "sliceRows": "(seqLength * numHeads if rank3 else seqLength) if rankOk else 0",
47
+ "sliceCount": "(batchSize if rank3 else batchSize * headBlocks) if rankOk else 0",
48
+ "sliceDispatchOk": "sliceCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
49
+ "workgroupSize": "max(1, min(tunables.WORKGROUP_SIZE, min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)))",
50
+ "laneUnits": "workgroupSize * itemsPerLane if rank3 else workgroupSize",
51
+ "quadHeadUnits": "headSize / 8 if headSize % 8 == 0 else 0",
52
+ "quadHeadStride": "headSize / 4 if headSize % 4 == 0 else 0",
53
+ "quadRotStart": "rotaryDim / 4 if rotaryDim % 4 == 0 else 0",
54
+ "quadRotUnits": "rotaryDim / 8 if rotaryDim % 8 == 0 else 0",
55
+ "quadHasTail": "quadHeadUnits > quadRotUnits",
56
+ "quadBlocks": "ceilDiv(sliceRows * quadHeadUnits, laneUnits)",
57
+ "quadOk": "headSize % 8 == 0 and rotaryDim % 8 == 0 and quadBlocks <= foldedDispatchCapacity and sliceDispatchOk",
58
+ "pairHeadUnits": "ceilDiv(headSize, 2)",
59
+ "pairHasTail": "pairHeadUnits > halfRotaryDim",
60
+ "pairTailGuard": "headSize % 2 == 1",
61
+ "pairBlocks": "ceilDiv(sliceRows * pairHeadUnits, laneUnits)",
62
+ "pairOk": "pairBlocks <= foldedDispatchCapacity and sliceDispatchOk",
63
+ "posOffsetOk": "ranks.positionIds <= 1 and numel(shapes.positionIds) == 1",
64
+ "posTableOk": "ranks.positionIds == 2 and dim(shapes.positionIds, 0) == batchSize and dim(shapes.positionIds, 1) == seqLength",
65
+ "interleaved": "attrs.interleaved != 0",
66
+ "scalarType": "dtypes.T",
67
+ "vectorType": "\"vec4<\" ~ dtypes.T ~ \">\""
68
+ },
69
+ "when": ["contractOk"],
70
+ "bindings": {
71
+ "x": { "elementType": "$vectorType" },
72
+ "position_ids": { "arg": "positionIds", "elementType": "$M" },
73
+ "cos_cache": { "arg": "cos", "elementType": "$vectorType" },
74
+ "sin_cache": { "arg": "sin", "elementType": "$vectorType" },
75
+ "y": { "elementType": "$vectorType" },
76
+ "x_scalar": { "arg": "x", "name": "x", "elementType": "$T" },
77
+ "cos_scalar": { "arg": "cos", "name": "cos_cache", "elementType": "$T" },
78
+ "sin_scalar": { "arg": "sin", "name": "sin_cache", "elementType": "$T" },
79
+ "y_scalar": { "arg": "y", "name": "y", "elementType": "$T" },
80
+ "params": { "struct": [{ "name": "sliceRows", "type": "u32", "value": "sliceRows" }] }
81
+ },
82
+ "variants": [
83
+ {
84
+ "id": "quad",
85
+ "priority": 10,
86
+ "when": ["posTableOk or posOffsetOk", "quadOk"],
87
+ "derive": {
88
+ "rank": "ranks.x",
89
+ "posOffset": "posOffsetOk",
90
+ "useVec4": "true",
91
+ "headUnits": "quadHeadUnits",
92
+ "headStride": "quadHeadStride",
93
+ "rotStart": "quadRotStart",
94
+ "rotUnits": "quadRotUnits",
95
+ "hasTail": "quadHasTail",
96
+ "tailGuard": "false",
97
+ "castIn": "\"vec4<f32>\"",
98
+ "castOut": "vectorType"
99
+ },
100
+ "passes": [
101
+ {
102
+ "id": "main",
103
+ "name": "RotaryEmbedding.Quad",
104
+ "shader": "rotary-embedding-slices.wgsl.jinja",
105
+ "bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params"],
106
+ "dispatch": {
107
+ "x": "min(quadBlocks, DISPATCH_FOLD_WIDTH)",
108
+ "y": "ceilDiv(quadBlocks, DISPATCH_FOLD_WIDTH)",
109
+ "z": "sliceCount"
110
+ }
111
+ }
112
+ ]
113
+ },
114
+ {
115
+ "id": "pair",
116
+ "priority": 0,
117
+ "when": ["posTableOk or posOffsetOk", "pairOk"],
118
+ "derive": {
119
+ "rank": "ranks.x",
120
+ "posOffset": "posOffsetOk",
121
+ "useVec4": "false",
122
+ "headUnits": "pairHeadUnits",
123
+ "headStride": "headSize",
124
+ "rotStart": "rotaryDim",
125
+ "rotUnits": "halfRotaryDim",
126
+ "hasTail": "pairHasTail",
127
+ "tailGuard": "pairTailGuard",
128
+ "castIn": "\"f32\"",
129
+ "castOut": "scalarType"
130
+ },
131
+ "passes": [
132
+ {
133
+ "id": "main",
134
+ "name": "RotaryEmbedding.Pair",
135
+ "shader": "rotary-embedding-slices.wgsl.jinja",
136
+ "bindings": ["x_scalar", "position_ids", "cos_scalar", "sin_scalar", "y_scalar", "params"],
137
+ "dispatch": {
138
+ "x": "min(pairBlocks, DISPATCH_FOLD_WIDTH)",
139
+ "y": "ceilDiv(pairBlocks, DISPATCH_FOLD_WIDTH)",
140
+ "z": "sliceCount"
141
+ }
142
+ }
143
+ ]
144
+ }
145
+ ]
146
+ }
build/webgpu/metadata.json ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "com.microsoft.RotaryEmbedding",
3
+ "id": "_com_microsoft_rotaryembedding_webgpu_91f411d",
4
+ "version": 1,
5
+ "license": "Apache-2.0",
6
+ "backend": { "type": "webgpu" },
7
+ "digest": {
8
+ "algorithm": "sha256",
9
+ "files": {
10
+ "bench.json": "8QUIGpk2EMC9wPHWw/ypeXIzxYOfJx9bfNsg8GpbBtc=",
11
+ "manifest.json": "6Gx395vXiLXEDd5ThN0DsKtZq+E2ttgZ8FPiGKkpaM4=",
12
+ "rotary-embedding-slices.wgsl.jinja": "Dl6jsIc9PxPRXaC8pn+mTys7XYeg6THauIZGbj8CmiQ=",
13
+ "test.json": "0DBcHqmXLgRmFXICRr3pcMCeOFxWKR5e3vQqzcnrv5E="
14
+ }
15
+ },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
+ "webgpu": {
18
+ "manifestSpec": "2.1",
19
+ "variants": { "quad": ["rotary-embedding-slices.wgsl.jinja"], "pair": ["rotary-embedding-slices.wgsl.jinja"] }
20
+ }
21
+ }
build/webgpu/rotary-embedding-slices.wgsl.jinja ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {{ env.wgsl.resourceDeclarations }}
2
+ // Rotary positional embedding. The dispatch pins one slice on z and folds that slice's own
3
+ // rotation groups across x and y, so every row of a slice shares one base offset and all the
4
+ // head geometry below is a compile-time constant: every division in the index math is by a
5
+ // literal rather than by a uniform.
6
+ //
7
+ // A slice is one batch at rank 3, where a token's heads are adjacent and a workgroup covers a
8
+ // whole token, so it reads each cos/sin row once. At rank 4 a token's heads are a sequence
9
+ // apart, so a slice is instead one batch and ITEMS consecutive heads: the invocation gathers
10
+ // its cos/sin row once and rotates the same rotation group in each of those heads.
11
+ const WG: u32 = {{ workgroupSize }}u;
12
+ const ITEMS: u32 = {{ itemsPerLane }}u;
13
+ // Rotation groups per head row, the row's stride, and the number of groups that rotate, all
14
+ // counted in the storage element the bindings use: a vec4 on the four-pair route and a scalar
15
+ // on the single-pair route.
16
+ const HEAD_UNITS: u32 = {{ headUnits }}u;
17
+ const HEAD_STRIDE: u32 = {{ headStride }}u;
18
+ const ROT_UNITS: u32 = {{ rotUnits }}u;
19
+ {% if hasTail %}
20
+ // First element past the rotary window.
21
+ const ROT_END: u32 = {{ rotStart }}u;
22
+ {% endif %}
23
+ const NUM_HEADS: u32 = {{ numHeads }}u;
24
+ {% if rank == 4 %}
25
+ const HEAD_BLOCKS: u32 = {{ headBlocks }}u;
26
+ {% endif %}
27
+ {% if hasTail %}
28
+
29
+ {% macro copy_tail(base) %}
30
+ let tail = ROT_END + (group - ROT_UNITS) * 2u;
31
+ y[{{ base }} + tail] = x[{{ base }} + tail];
32
+ {% if tailGuard %}
33
+ // An odd head size leaves the last group with a single element to copy.
34
+ if (tail + 1u < HEAD_STRIDE) {
35
+ y[{{ base }} + tail + 1u] = x[{{ base }} + tail + 1u];
36
+ }
37
+ {% else %}
38
+ y[{{ base }} + tail + 1u] = x[{{ base }} + tail + 1u];
39
+ {% endif %}
40
+ {% endmacro %}
41
+ {% endif %}
42
+
43
+ {% macro rotate(base) %}
44
+ {% if interleaved and useVec4 %}
45
+ // Four interleaved pairs span two vec4s: (w0.x, w0.y), (w0.z, w0.w), (w1.x, w1.y),
46
+ // (w1.z, w1.w). Deinterleave, rotate, and reinterleave on the way back out.
47
+ let pairBase = {{ base }} + group * 2u;
48
+ let w0 = vec4<f32>(x[pairBase]);
49
+ let w1 = vec4<f32>(x[pairBase + 1u]);
50
+ let a = vec4<f32>(w0.x, w0.z, w1.x, w1.z);
51
+ let b = vec4<f32>(w0.y, w0.w, w1.y, w1.w);
52
+ let ra = a * cf - b * sf;
53
+ let rb = a * sf + b * cf;
54
+ y[pairBase] = {{ castOut }}(vec4<f32>(ra.x, rb.x, ra.y, rb.y));
55
+ y[pairBase + 1u] = {{ castOut }}(vec4<f32>(ra.z, rb.z, ra.w, rb.w));
56
+ {% elif interleaved %}
57
+ // Adjacent elements form the rotated pair.
58
+ let pairBase = {{ base }} + group * 2u;
59
+ let a = {{ castIn }}(x[pairBase]);
60
+ let b = {{ castIn }}(x[pairBase + 1u]);
61
+ y[pairBase] = {{ castOut }}(a * cf - b * sf);
62
+ y[pairBase + 1u] = {{ castOut }}(a * sf + b * cf);
63
+ {% else %}
64
+ // Split halves: ROT_UNITS is both the rotating group count and the offset from an
65
+ // element of the first half to its partner in the second.
66
+ let a = {{ castIn }}(x[{{ base }} + group]);
67
+ let b = {{ castIn }}(x[{{ base }} + ROT_UNITS + group]);
68
+ y[{{ base }} + group] = {{ castOut }}(a * cf - b * sf);
69
+ y[{{ base }} + ROT_UNITS + group] = {{ castOut }}(a * sf + b * cf);
70
+ {% endif %}
71
+ {% endmacro %}
72
+
73
+ @compute @workgroup_size(WG)
74
+ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
75
+ {% if rank == 3 %}
76
+ // wg.y carries the high bits of the group index past the per-axis dispatch fold width.
77
+ let slice = wg.z;
78
+ let sliceUnits = params.sliceRows * HEAD_UNITS;
79
+ let sliceBase = slice * params.sliceRows * HEAD_STRIDE;
80
+ let laneUnit = (wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u) * (WG * ITEMS) + lid.x;
81
+ {% for item in range(itemsPerLane) %}
82
+ {
83
+ let unit = laneUnit + {{ item }}u * WG;
84
+ if (unit < sliceUnits) {
85
+ let group = unit % HEAD_UNITS;
86
+ let row = unit / HEAD_UNITS;
87
+ let base = sliceBase + row * HEAD_STRIDE;
88
+ let token = row / NUM_HEADS;
89
+ {% if hasTail %}
90
+ if (group >= ROT_UNITS) {
91
+ {{ copy_tail("base") }}
92
+ } else {
93
+ {% endif %}
94
+ {% if posOffset %}
95
+ // One base offset for the whole request: token s reads cache row p[0] + s.
96
+ let position = position_ids[0] + token;
97
+ {% else %}
98
+ // The (batch_size, sequence_length) table. params.sliceRows / NUM_HEADS is the sequence
99
+ // length, so this slice's table row is slice * sequence_length + token.
100
+ let position = position_ids[slice * (params.sliceRows / NUM_HEADS) + token];
101
+ {% endif %}
102
+ let cache = position * ROT_UNITS + group;
103
+ let cf = {{ castIn }}(cos_cache[cache]);
104
+ let sf = {{ castIn }}(sin_cache[cache]);
105
+ {{ rotate("base") }}
106
+ {% if hasTail %}
107
+ }
108
+ {% endif %}
109
+ }
110
+ }
111
+ {% endfor %}
112
+ {% else %}
113
+ // A rank-4 slice is one batch and ITEMS consecutive heads; wg.y carries the high bits of
114
+ // the group index past the per-axis dispatch fold width.
115
+ let slice = wg.z;
116
+ let headBase = slice % HEAD_BLOCKS * ITEMS;
117
+ let batch = slice / HEAD_BLOCKS;
118
+ let sliceUnits = params.sliceRows * HEAD_UNITS;
119
+ let unit = (wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u) * WG + lid.x;
120
+ if (unit >= sliceUnits) {
121
+ return;
122
+ }
123
+ let group = unit % HEAD_UNITS;
124
+ let token = unit / HEAD_UNITS;
125
+ let rowBase = (batch * NUM_HEADS + headBase) * params.sliceRows + token;
126
+ {% if hasTail %}
127
+ if (group >= ROT_UNITS) {
128
+ {% for item in range(itemsPerLane) %}
129
+ {
130
+ let base = (rowBase + {{ item }}u * params.sliceRows) * HEAD_STRIDE;
131
+ {{ copy_tail("base") }}
132
+ }
133
+ {% endfor %}
134
+ return;
135
+ }
136
+ {% endif %}
137
+ {% if posOffset %}
138
+ // One base offset for the whole request: token s reads cache row p[0] + s.
139
+ let position = position_ids[0] + token;
140
+ {% else %}
141
+ // The (batch_size, sequence_length) table; params.sliceRows is the sequence length.
142
+ let position = position_ids[batch * params.sliceRows + token];
143
+ {% endif %}
144
+ // Gathered once for the whole head block: cos and sin depend on the token and the rotation
145
+ // group, never on the head.
146
+ let cache = position * ROT_UNITS + group;
147
+ let cf = {{ castIn }}(cos_cache[cache]);
148
+ let sf = {{ castIn }}(sin_cache[cache]);
149
+ {% for item in range(itemsPerLane) %}
150
+ {
151
+ let base = (rowBase + {{ item }}u * params.sliceRows) * HEAD_STRIDE;
152
+ {{ rotate("base") }}
153
+ }
154
+ {% endfor %}
155
+ {% endif %}
156
+ }
build/webgpu/test.json ADDED
@@ -0,0 +1,1181 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fixtureArrays": {
3
+ "ort_contrib_interleaved_small_llama_offset_f32_input_x": [-1.0408, 0.9166, -1.3042, -1.1097, -0.132, -0.2751, -0.235, 0.0937, -1.2188, 1.1676, -1.0574, -0.1188, -0.7396, -1.2425, -0.1752, 0.699, -0.811, 0.6737, -1.1233, -0.0919, -0.6861, 0.7202, 0.1963, 0.6142],
4
+ "ort_contrib_interleaved_small_llama_offset_f32_input_cos": [1, 1, 0.5403, 0.9999, -0.4161, 0.9998, -0.99, 0.9996, -0.6536, 0.9992, 0.2837, 0.9988, 0.9602, 0.9982, 0.7539, 0.9976],
5
+ "ort_contrib_interleaved_small_llama_offset_f32_input_sin": [0, 0, 0.8415, 0.01, 0.9093, 0.02, 0.1411, 0.03, -0.7568, 0.04, -0.9589, 0.05, -0.2794, 0.06, 0.657, 0.0699],
6
+ "ort_contrib_interleaved_small_llama_offset_f32_output_y": [-1.0408, 0.9166, -1.3042, -1.1097, -0.132, -0.2751, -0.235, 0.0937, -1.6411, -0.3948, -1.0561, -0.1294, 0.646, -1.2937, -0.1822, 0.6972, -0.2751, -1.0178, -1.1212, -0.1143, -0.3694, -0.9235, 0.184, 0.618],
7
+ "ort_contrib_not_interleaved_small_llama_table_f32_input_x": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, 1.0076, -0.7529, -0.225, -0.4327, -1.5071, -0.4586, -0.8663, -0.2656, 0.1665, 0.7911, -0.932, -0.8579, -1.0574, -0.1188, -0.9078, 0.3452, -0.5713, -0.2351, -0.848, 0.5266, -1.2944, -0.0243, -0.2354, -0.7087, -0.9647, -0.0991, -0.2994, -0.065, -1.572, -1.3211],
8
+ "ort_contrib_not_interleaved_small_llama_table_f32_output_y": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, 1.0076, -0.7529, -0.225, -0.4327, -1.5071, -0.4586, -0.8663, -0.2656, 0.1665, 0.7911, -0.932, -0.8579, -0.8618, -0.0922, -0.9073, -0.7032, -0.5762, -0.2371, -0.4377, 0.537, -1.2929, -0.7267, -0.2107, -0.7115, -0.4666, -0.0261, -0.2965, -0.8469, -1.5749, -1.3217],
9
+ "ort_contrib_interleaved_large_llama_offset_f32_input_x": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, -1.019, 0.3157, -1.6036, 1.8493, 0.0447, 1.5853, 0.1036, -0.3514, 0.2421, 0.6463, 0.873, -0.9276, 1.0311, -1.9557, -0.1482, 1.7376, 2.2039, -0.6589, -1.0574, -0.1188, -0.9078, 0.3452, -0.5713, -0.2351, -0.5912, 1.1312, 0.7562, -1.2023, -0.5833, -0.4407, 0.1766, 1.0224, -0.4826, -0.5421, -0.5342, -0.6413, 1.3314, -0.4498, 0.5493, 0.0539, 0.2601, 0.857, 1.0076, -0.7529, -0.225, -0.4327, -1.5071, -0.4586, -1.9791, 0.7787, -0.7749, -0.1398, 1.1414, -0.6354, 0.0352, -0.4765, -0.0409, 1.1993, 0.5374, -0.193, 2.5211, -0.0452, -0.3105, -0.9407, -0.0034, 1.5199, -0.848, 0.5266, 0.0299, -0.0498, 1.0651, 0.886, -1.4702, -0.2134, -0.8707, 1.6159, -0.2356, 0.9444, 0.5937, 0.7203, 0.5061, 1.5192, -0.4897, 0.9231, 0.2654, -0.1441, 0.5407, -1.5476, 0.6455, -1.1382, 0.464, -0.4986, 0.1289, 2.7631, 0.1405, 1.1191, 2.1134, -0.9754, 0.1757, -0.1319, -0.2735, 0.3355, -0.6008, -1.1164, 0.2577, -0.7226, -0.9244, 1.8737, 0.6052, 1.1904, 1.2195, -0.047, -1.0914, 1.0223, 0.3152, 1.7528, -0.765, 1.8299, -0.2784, -0.2719, 0.1885, 2.1432, 0.8527, 0.0965, -0.0625, 0.8269, 1.0122, -1.4482, -0.0644, 0.3215, 0.5908, -1.4197, 0.2113, 0.0306, 0.3604, 0.3166, -0.8975, -0.6393, -1.2944, -0.0243, -0.2354, -0.7087, 1.1566, 0.4296, 0.5599, -0.7776, 0.3339, 0.1759, 2.1108, 1.0702, 0.8279, -0.2969, 0.712, -0.2068, -0.1548, 0.1553, 0.6207, -0.169, -0.5816, 1.2632, 0.0695, 1.1862, -1.1874, -0.7468, -0.932, -0.8579, -0.9647, -0.0991, 0.0195, 1.1213, -1.4873, -0.2043, -1.0466, -1.5772, -0.0489, 0.343, 0.1264, 0.1519, -1.3639, -1.6593, 1.8127, -1.4459, -0.2158, -0.9792, -1.4392, 0.6508, 0.8964, 0.5717, -0.239, 0.6983, -1.3416, 0.2715, -0.2852, 0.6051, 0.2167, -0.2181, -1.6306, 1.4788, 0.2754, -0.0261, -0.4618, -0.5646, -1.0389, 0.5819, 1.3697, 0.0002, 1.5333, -1.0556, -0.1254, 0.1527, -0.5996, -1.0962, 1.6327, 1.3951, 0.8784, 0.3389, 1.2907, 0.3124, 0.7299, 1.422, 0.3375, 0.0438, 1.8698, -0.2635, -2.0799, -0.6313, 0.409, -1.1458, 0.0784, -1.8848, -1.6165, 0.6179, 0.9905, -0.0729, 0.5054, -0.6681, -1.4382, 1.7547, -0.9605, -0.4558, -1.6105, 0.2979, 1.1537, -1.5604, 1.2779, -1.2514, 0.6056, 0.5763, -3.3558, 0.2836, 0.6909, -0.7631, 2.4451, -0.35, 1.3289, -0.6494, 0.3478, 1.0038, -0.2937, 0.9238, -1.2185, 0.4138, 0.5033, 0.9174, 1.8131, 1.4436, -0.4207, 0.022, -0.6807, -1.3306, 1.5646, 0.3338, 0.7105, 0.4683, -0.6179, 0.0818, -0.0488, -0.981, -1.3632, 0.0929, -1.7926, -0.2921, -0.4792, 0.6756, -0.3413, -0.2242, -0.2111, 0.6282, 0.1667, -1.4055, 1.5895, 1.0838, -0.9077, -0.806, 0.7967, -2.9351, 2.4179, -0.4026, 0.6451, 1.6845, -0.0901, 0.6106, 2.3603, 1.3908, -0.7917, -0.6734, -0.1213, -1.1116, -0.7401, -0.7879, 0.0606, -2.3337, -1.2603, -1.7245, -0.3533, -0.9421, -0.1776, 0.3992, -1.7142, -0.5319, -0.8848, 0.6513, 1.0002, -1.4699, -1.4254, 0.7013, 0.2414, 0.2551, -0.7457, 0.3133, -1.0941, -0.3682, -0.0163, -0.0645, -0.8101, 0.1415, 0.0551, 0.5873, -0.5887, -1.4733, -0.8565, 0.74, -0.5033, 0.0553, 0.9265, -0.8652, -0.0288, -0.2209, 0.061, 0.6776, 0.4361, -0.8052, 0.3955, 0.8988, 0.8238, 0.2262, 1.2912, 0.6488, 1.2114, 1.3569, 0.2983, 0.4718, -1.1936, 0.7928, -0.8665, 0.9468, 1.1629, 0.0616, -1.3136, -0.2764, 0.0277, -0.1126, 0.2342, -0.5866, -1.8219, 1.1079, 0.5795, -1.4249],
10
+ "ort_contrib_interleaved_large_llama_offset_f32_input_cos": [1, 1, 1, 0.5403, 0.9989, 1, -0.4161, 0.9957, 1, -0.99, 0.9903, 1, -0.6536, 0.9828, 1, 0.2837, 0.9732, 0.9999, 0.9602, 0.9615, 0.9999, 0.7539, 0.9477, 0.9999, -0.1455, 0.9318, 0.9999, -0.9111, 0.914, 0.9998, -0.8391, 0.8942, 0.9998, 0.0044, 0.8725, 0.9997, 0.8439, 0.8488, 0.9997, 0.9074, 0.8234, 0.9996, 0.1367, 0.7962, 0.9995, -0.7597, 0.7673, 0.9995],
11
+ "ort_contrib_interleaved_large_llama_offset_f32_input_sin": [0, 0, 0, 0.8415, 0.0464, 0.0022, 0.9093, 0.0927, 0.0043, 0.1411, 0.1388, 0.0065, -0.7568, 0.1846, 0.0086, -0.9589, 0.23, 0.0108, -0.2794, 0.2749, 0.0129, 0.657, 0.3192, 0.0151, 0.9894, 0.3629, 0.0172, 0.4121, 0.4057, 0.0194, -0.544, 0.4477, 0.0215, -1, 0.4887, 0.0237, -0.5366, 0.5286, 0.0259, 0.4202, 0.5675, 0.028, 0.9906, 0.605, 0.0302, 0.6503, 0.6413, 0.0323]
12
+ },
13
+ "cases": [
14
+ {
15
+ "name": "ort_contrib_interleaved_small_llama_offset_f32",
16
+ "attrs": { "interleaved": 1 },
17
+ "provenance": {
18
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
19
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_Interleaved_SmallData_LlamaMSFT",
20
+ "notes": "Format-0 position ids (a one-element base offset) with num_heads omitted, so head_size is 2 * cos_cache width and the head count is inferred from hidden_size."
21
+ },
22
+ "inputs": {
23
+ "x": {
24
+ "dtype": "float32",
25
+ "shape": [1, 3, 8],
26
+ "data": {
27
+ "kind": "values",
28
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_x" }
29
+ }
30
+ },
31
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [0] } },
32
+ "cos": {
33
+ "dtype": "float32",
34
+ "shape": [8, 2],
35
+ "data": {
36
+ "kind": "values",
37
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_cos" }
38
+ }
39
+ },
40
+ "sin": {
41
+ "dtype": "float32",
42
+ "shape": [8, 2],
43
+ "data": {
44
+ "kind": "values",
45
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_sin" }
46
+ }
47
+ }
48
+ },
49
+ "outputs": {
50
+ "y": {
51
+ "dtype": "float32",
52
+ "shape": [1, 3, 8],
53
+ "data": {
54
+ "kind": "values",
55
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_output_y" }
56
+ }
57
+ }
58
+ },
59
+ "tolerance": 0.002,
60
+ "relTolerance": 0.0001
61
+ },
62
+ {
63
+ "name": "ort_contrib_interleaved_small_llama_offset_f16",
64
+ "attrs": { "interleaved": 1 },
65
+ "provenance": {
66
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
67
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_Interleaved_SmallData_LlamaMSFT",
68
+ "notes": "The float16 arm of the same upstream case; ORT runs every RunTests case at float32 and float16 against the same expected values with a 0.002 absolute tolerance."
69
+ },
70
+ "inputs": {
71
+ "x": {
72
+ "dtype": "float16",
73
+ "shape": [1, 3, 8],
74
+ "data": {
75
+ "kind": "values",
76
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_x" }
77
+ }
78
+ },
79
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [0] } },
80
+ "cos": {
81
+ "dtype": "float16",
82
+ "shape": [8, 2],
83
+ "data": {
84
+ "kind": "values",
85
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_cos" }
86
+ }
87
+ },
88
+ "sin": {
89
+ "dtype": "float16",
90
+ "shape": [8, 2],
91
+ "data": {
92
+ "kind": "values",
93
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_sin" }
94
+ }
95
+ }
96
+ },
97
+ "outputs": {
98
+ "y": {
99
+ "dtype": "float16",
100
+ "shape": [1, 3, 8],
101
+ "data": {
102
+ "kind": "values",
103
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_output_y" }
104
+ }
105
+ }
106
+ },
107
+ "tolerance": 0.002,
108
+ "relTolerance": 0.002
109
+ },
110
+ {
111
+ "name": "ort_contrib_not_interleaved_small_llama_table_f32",
112
+ "attrs": { "interleaved": 0 },
113
+ "provenance": {
114
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
115
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_NotInterleaved_SmallData_LlamaMSFT",
116
+ "notes": "Format-1 position ids over three heads of six, with num_heads omitted so the head size comes from the cache width."
117
+ },
118
+ "inputs": {
119
+ "x": {
120
+ "dtype": "float32",
121
+ "shape": [1, 2, 18],
122
+ "data": {
123
+ "kind": "values",
124
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_not_interleaved_small_llama_table_f32_input_x" }
125
+ }
126
+ },
127
+ "positionIds": { "dtype": "uint32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
128
+ "cos": {
129
+ "dtype": "float32",
130
+ "shape": [4, 3],
131
+ "data": {
132
+ "kind": "values",
133
+ "values": [1.0, 1.0, 1.0, 0.5403, 0.9989, 1.0, -0.4161, 0.9957, 1.0, -0.99, 0.9903, 1.0]
134
+ }
135
+ },
136
+ "sin": {
137
+ "dtype": "float32",
138
+ "shape": [4, 3],
139
+ "data": {
140
+ "kind": "values",
141
+ "values": [0.0, 0.0, 0.0, 0.8415, 0.0464, 0.0022, 0.9093, 0.0927, 0.0043, 0.1411, 0.1388, 0.0065]
142
+ }
143
+ }
144
+ },
145
+ "outputs": {
146
+ "y": {
147
+ "dtype": "float32",
148
+ "shape": [1, 2, 18],
149
+ "data": {
150
+ "kind": "values",
151
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_not_interleaved_small_llama_table_f32_output_y" }
152
+ }
153
+ }
154
+ },
155
+ "tolerance": 0.002,
156
+ "relTolerance": 0.0001
157
+ },
158
+ {
159
+ "name": "ort_contrib_not_interleaved_small_llama_table_f16",
160
+ "attrs": { "interleaved": 0 },
161
+ "provenance": {
162
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
163
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_NotInterleaved_SmallData_LlamaMSFT",
164
+ "notes": "The float16 arm of the same upstream case."
165
+ },
166
+ "inputs": {
167
+ "x": {
168
+ "dtype": "float16",
169
+ "shape": [1, 2, 18],
170
+ "data": {
171
+ "kind": "values",
172
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_not_interleaved_small_llama_table_f32_input_x" }
173
+ }
174
+ },
175
+ "positionIds": { "dtype": "uint32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
176
+ "cos": {
177
+ "dtype": "float16",
178
+ "shape": [4, 3],
179
+ "data": {
180
+ "kind": "values",
181
+ "values": [1.0, 1.0, 1.0, 0.5403, 0.9989, 1.0, -0.4161, 0.9957, 1.0, -0.99, 0.9903, 1.0]
182
+ }
183
+ },
184
+ "sin": {
185
+ "dtype": "float16",
186
+ "shape": [4, 3],
187
+ "data": {
188
+ "kind": "values",
189
+ "values": [0.0, 0.0, 0.0, 0.8415, 0.0464, 0.0022, 0.9093, 0.0927, 0.0043, 0.1411, 0.1388, 0.0065]
190
+ }
191
+ }
192
+ },
193
+ "outputs": {
194
+ "y": {
195
+ "dtype": "float16",
196
+ "shape": [1, 2, 18],
197
+ "data": {
198
+ "kind": "values",
199
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_not_interleaved_small_llama_table_f32_output_y" }
200
+ }
201
+ }
202
+ },
203
+ "tolerance": 0.002,
204
+ "relTolerance": 0.002
205
+ },
206
+ {
207
+ "name": "ort_contrib_interleaved_large_llama_offset_f32",
208
+ "attrs": { "interleaved": 1 },
209
+ "provenance": {
210
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
211
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_Interleaved_LargeData_LlamaMSFT",
212
+ "notes": "Two batches of eight tokens over four heads of six, gathered through a format-0 base offset of zero."
213
+ },
214
+ "inputs": {
215
+ "x": {
216
+ "dtype": "float32",
217
+ "shape": [2, 8, 24],
218
+ "data": {
219
+ "kind": "values",
220
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_x" }
221
+ }
222
+ },
223
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [0] } },
224
+ "cos": {
225
+ "dtype": "float32",
226
+ "shape": [16, 3],
227
+ "data": {
228
+ "kind": "values",
229
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_cos" }
230
+ }
231
+ },
232
+ "sin": {
233
+ "dtype": "float32",
234
+ "shape": [16, 3],
235
+ "data": {
236
+ "kind": "values",
237
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_sin" }
238
+ }
239
+ }
240
+ },
241
+ "outputs": {
242
+ "y": {
243
+ "dtype": "float32",
244
+ "shape": [2, 8, 24],
245
+ "data": {
246
+ "kind": "values",
247
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, -1.019, 0.3157, -1.6036, 1.8493, 0.0447, 1.5853, 0.1036, -0.3514, 0.2421, 0.6463, 0.873, -0.9276, 1.0311, -1.9557, -0.1482, 1.7376, 2.2039, -0.6589, -0.4713, -0.954, -0.9229, 0.3027, -0.5708, -0.2363, -1.2713, 0.1137, 0.8112, -1.1659, -0.5824, -0.4419, -0.7649, 0.7011, -0.4569, -0.5639, -0.5328, -0.6424, 1.0979, 0.8773, 0.5462, 0.0793, 0.2582, 0.8576, 0.2653, 1.2295, -0.1839, -0.4517, -1.5052, -0.4651, 0.1155, -2.1237, -0.7586, -0.211, 1.1441, -0.6304, 0.4186, 0.2303, -0.1519, 1.1903, 0.5382, -0.1906, -1.008, 2.3112, -0.222, -0.9655, -0.0099, 1.5198, 0.7652, -0.641, 0.0365, -0.0452, 1.0593, 0.8929, 1.4856, 0.0038, -1.0865, 1.4794, -0.2417, 0.9428, -0.6894, -0.6293, 0.2904, 1.5747, -0.4956, 0.9199, -0.2424, 0.1801, 0.7503, -1.4576, 0.6529, -1.134, -0.6807, -0.0252, -0.3834, 2.7394, 0.1308, 1.1203, -2.1196, -0.9618, 0.197, -0.0972, -0.2764, 0.3332, -0.4522, 1.1844, 0.3867, -0.6626, -0.9405, 1.8656, 0.5053, -1.2361, 1.2072, 0.1789, -1.1002, 1.0129, 1.7702, 0.1949, -1.1653, 1.6049, -0.2755, -0.2749, 2.1087, 0.4272, 0.8076, 0.29, -0.0714, 0.8261, -1.1016, -1.3814, -0.1366, 0.2981, 0.606, -1.4132, 0.0893, -0.1939, 0.2779, 0.391, -0.8906, -0.6489, -1.2496, 0.3383, -0.0315, -0.7461, 1.151, 0.4445, 0.3203, -0.9031, 0.2727, 0.2609, 2.0968, 1.0974, 0.712, -0.5164, 0.7415, -0.0031, -0.1568, 0.1533, 0.5487, -0.3357, -0.9064, 1.0546, 0.0542, 1.187, -0.4045, -1.3431, -0.6094, -1.1105, -0.9631, -0.1137, -0.7219, 0.8582, -1.3443, -0.6684, -1.0227, -1.5929, -0.2622, 0.2264, 0.0713, 0.1843, -1.3387, -1.6797, 2.3165, 0.1009, 0.1081, -0.9969, -1.4488, 0.6291, 0.8964, 0.5717, -0.239, 0.6983, -1.3416, 0.2715, -0.2852, 0.6051, 0.2167, -0.2181, -1.6306, 1.4788, 0.2754, -0.0261, -0.4618, -0.5646, -1.0389, 0.5819, 1.3697, 0.0002, 1.5333, -1.0556, -0.1254, 0.1527, 0.5985, -1.0968, 1.5662, 1.4693, 0.8776, 0.3408, 0.4345, 1.2549, 0.6631, 1.4543, 0.3374, 0.0445, 1.232, 1.4311, -2.0483, -0.7272, 0.4114, -1.1449, 1.6283, -0.9524, -1.6435, 0.5422, 0.9907, -0.0708, 0.3972, 0.7376, -1.5947, 1.6138, -0.9586, -0.46, 0.3993, -1.5884, 1.2934, -1.4467, 1.2833, -1.2459, -0.776, 0.3108, -3.3677, -0.0287, 0.6942, -0.7601, -0.6993, 2.369, 1.3834, -0.5234, 0.3435, 1.0053, 0.1604, -0.956, -1.2641, 0.2406, 0.4973, 0.9206, -1.9987, -1.1733, -0.4197, -0.0366, -0.672, -1.335, -1.596, -0.1097, 0.6386, 0.5624, -0.6184, 0.0778, 0.1867, 0.9643, -1.3629, -0.0972, -1.7907, -0.3037, 0.8245, -0.0789, -0.294, -0.2833, -0.2165, 0.6264, -1.1726, 0.7926, 1.3621, 1.3586, -0.9007, -0.8138, -2.7421, 1.3155, 2.4507, 0.0507, 0.6305, 1.69, 0.521, -0.3309, 2.063, 1.8026, -0.7859, -0.6802, -1.1003, -0.199, -0.5391, -0.937, 0.0857, -2.333, -2.0112, 0.7193, -0.1272, -0.9981, -0.1818, 0.3973, -0.9963, 1.4929, -1.0109, 0.4304, 1.016, -1.459, 0.2682, 1.5658, 0.1762, 0.3038, -0.7491, 0.3052, -1.1534, -0.0478, 0.0021, -0.0665, -0.8118, 0.131, 0.2171, 0.5485, -0.161, -1.5784, -0.866, 0.7289, -0.4678, 0.1937, 1.1287, -0.5772, -0.0259, -0.2212, 0.2479, 0.6336, 0.6407, -0.6543, 0.3838, 0.9039, 0.4724, 0.7117, 1.0165, 1.027, 1.1908, 1.375, -0.085, 0.5517, -1.3842, 0.3703, -0.8806, 0.9336, 0.8362, 0.8105, -1.1566, -0.6813, 0.0294, -0.1122, 0.562, -0.2884, -2.0803, 0.4684, 0.6009, -1.416]
248
+ }
249
+ }
250
+ },
251
+ "tolerance": 0.002,
252
+ "relTolerance": 0.0001
253
+ },
254
+ {
255
+ "name": "ort_contrib_not_interleaved_large_llama_offset_f32",
256
+ "attrs": { "interleaved": 0 },
257
+ "provenance": {
258
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
259
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_NotInterleaved_LargeData_LlamaMSFT",
260
+ "notes": "The split-halves pairing over the same two-batch geometry and a format-0 base offset."
261
+ },
262
+ "inputs": {
263
+ "x": {
264
+ "dtype": "float32",
265
+ "shape": [2, 8, 24],
266
+ "data": {
267
+ "kind": "values",
268
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_x" }
269
+ }
270
+ },
271
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [0] } },
272
+ "cos": {
273
+ "dtype": "float32",
274
+ "shape": [16, 3],
275
+ "data": {
276
+ "kind": "values",
277
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_cos" }
278
+ }
279
+ },
280
+ "sin": {
281
+ "dtype": "float32",
282
+ "shape": [16, 3],
283
+ "data": {
284
+ "kind": "values",
285
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_sin" }
286
+ }
287
+ }
288
+ },
289
+ "outputs": {
290
+ "y": {
291
+ "dtype": "float32",
292
+ "shape": [2, 8, 24],
293
+ "data": {
294
+ "kind": "values",
295
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, -1.019, 0.3157, -1.6036, 1.8493, 0.0447, 1.5853, 0.1036, -0.3514, 0.2421, 0.6463, 0.873, -0.9276, 1.0311, -1.9557, -0.1482, 1.7376, 2.2039, -0.6589, -0.8618, -0.0922, -0.9073, -0.7032, -0.5762, -0.2371, 0.6923, 1.1571, 0.7572, -1.1471, -0.5302, -0.4391, 0.5516, 1.0461, -0.4812, -0.1443, -0.4862, -0.6423, 0.674, -0.4614, 0.5475, 1.1495, 0.2389, 0.8582, -0.0259, -0.6099, -0.223, 1.0963, -1.5704, -0.4595, 0.9507, 0.6696, -0.7721, -1.7415, 1.2087, -0.6387, -1.1052, -0.5243, -0.04, -0.4671, 0.4909, -0.1931, -0.1937, -0.0447, -0.3171, 2.6839, -0.0076, 1.5185, 0.8465, 0.3737, 0.0242, -0.0703, 1.1279, 0.8862, 1.2275, -0.1786, -0.8767, -1.8072, -0.263, 0.9387, -0.8021, 0.7813, 0.5001, -1.4202, -0.385, 0.9263, -0.0443, -0.2323, 0.548, 1.5696, 0.6193, -1.1346, 1.7878, -0.516, 0.1192, -2.1572, 0.046, 1.1202, -1.4812, -0.9082, 0.1728, -1.5132, -0.4489, 0.337, -0.1541, -0.9266, 0.2416, 0.927, -1.1146, 1.8758, -0.4312, 1.3714, 1.2106, -0.4272, -0.8529, 1.0328, 1.8441, 1.7698, -0.762, 0.2168, 0.1322, -0.2802, 0.146, 2.1002, 0.8437, -0.1534, 0.4321, 0.836, 0.5955, -1.5452, -0.0491, -0.8794, 0.2418, -1.4203, 0.3635, 0.2362, 0.3672, -0.1128, -0.8664, -0.6354, -1.4409, -0.3413, -0.2409, -0.3188, 1.1054, 0.4265, 0.5867, -1.3279, 0.3201, 0.0125, 1.8157, 1.0745, 0.7372, -0.2429, 0.71, -0.4299, -0.2304, 0.1645, 0.9489, -0.1816, -0.5968, 1.0394, 0.0204, 1.1786, -0.3315, -0.3997, -0.9304, -1.4268, -1.1526, -0.1132, 0.149, 1.3967, -1.4634, -0.1412, -0.6339, -1.5995, -0.1366, 0.7604, 0.1514, 0.0824, -1.183, -1.6572, 2.0099, -0.9108, -0.2256, 0.4527, -1.8254, 0.6475, 0.8964, 0.5717, -0.239, 0.6983, -1.3416, 0.2715, -0.2852, 0.6051, 0.2167, -0.2181, -1.6306, 1.4788, 0.2754, -0.0261, -0.4618, -0.5646, -1.0389, 0.5819, 1.3697, 0.0002, 1.5333, -1.0556, -0.1254, 0.1527, -1.4979, -1.1358, 1.632, 0.2493, 0.8266, 0.3424, -0.4992, 0.2964, 0.7298, 1.8544, 0.3516, 0.0454, 1.5415, -0.2822, -2.0774, 1.2323, 0.3963, -1.1503, -0.4775, -1.9287, -1.6164, 0.3998, 0.902, -0.0764, -1.8059, -0.5762, -1.4362, -0.2706, -1.0183, -0.462, 2.0891, 0.1782, 1.1591, -0.8151, 1.3, -1.2464, -0.5099, 0.5098, -3.3525, 0.4326, 0.7414, -0.7775, -0.4271, -0.3807, 1.3245, 2.4936, 0.3139, 1.0095, 0.2323, 0.845, -1.2244, -0.4511, 0.6266, 0.9095, -1.7981, 1.5241, -0.4121, 0.2341, -0.4737, -1.3333, -1.615, 0.4164, 0.71, -0.2429, -0.5656, 0.0863, 0.0352, -0.7227, -1.3613, -0.0988, -1.9114, -0.3009, 0.1435, 0.7029, -0.3467, 0.5092, -0.0828, 0.6253, 0.7113, -1.2138, 1.5964, -0.8346, -1.1515, -0.7923, -0.8254, -3.0038, 2.4033, -0.3398, 0.0922, 1.7053, 1.1114, 0.7462, 2.366, -0.8409, -0.6654, -0.653, -0.7899, -1.0957, -0.7149, -0.1072, -0.1967, -2.3416, -1.2609, -1.6375, -0.3576, 0.9413, -0.5694, 0.3954, 0.1383, -0.7477, -0.8689, 1.8286, 0.851, -1.4793, -0.1597, 0.8541, 0.238, 1.4392, -0.5644, 0.3158, -1.0686, -0.1313, -0.0181, 0.2438, -0.8801, 0.1413, -0.3587, 0.8002, -0.5982, -1.4301, -0.662, 0.7324, -0.725, 0.061, 0.9293, -0.6902, -0.0125, -0.2089, -0.1664, 0.5428, 0.4245, -0.7901, 0.5665, 0.9044, 0.1948, -0.1723, 1.2705, 1.0303, 1.2202, 1.3762, -0.2959, 0.7237, -1.2077, 0.7937, -0.6705, 0.9287, 1.0583, 0.0496, -1.3118, 0.5556, 0.0459, -0.1324, -0.5513, -0.7409, -1.8002, 0.9892, 0.3619, -1.4522]
296
+ }
297
+ }
298
+ },
299
+ "tolerance": 0.002,
300
+ "relTolerance": 0.0001
301
+ },
302
+ {
303
+ "name": "ort_contrib_custom_rotary_dim_phi_table_f32",
304
+ "attrs": { "interleaved": 0, "rotary_embedding_dim": 4, "num_heads": 1 },
305
+ "provenance": {
306
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
307
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_CustomRotaryDim_SmallData_Phi",
308
+ "notes": "Partial rotation: four of six head elements rotate and the two-element tail is copied unchanged."
309
+ },
310
+ "inputs": {
311
+ "x": {
312
+ "dtype": "float32",
313
+ "shape": [1, 2, 6],
314
+ "data": {
315
+ "kind": "values",
316
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, 1.0076, -0.7529, -0.225, -0.4327, -1.5071, -0.4586]
317
+ }
318
+ },
319
+ "positionIds": { "dtype": "uint32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
320
+ "cos": { "dtype": "float32", "shape": [2, 2], "data": { "kind": "values", "values": [1.0, 1.0, 1.0, 0.5403] } },
321
+ "sin": { "dtype": "float32", "shape": [2, 2], "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.8415] } }
322
+ },
323
+ "outputs": {
324
+ "y": {
325
+ "dtype": "float32",
326
+ "shape": [1, 2, 6],
327
+ "data": {
328
+ "kind": "values",
329
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, 1.0076, -0.0427, -0.225, -0.8673, -1.5071, -0.4586]
330
+ }
331
+ }
332
+ },
333
+ "tolerance": 0.002,
334
+ "relTolerance": 0.0001
335
+ },
336
+ {
337
+ "name": "ort_contrib_custom_rotary_dim_phi_packed_batching_f32",
338
+ "attrs": { "interleaved": 0, "rotary_embedding_dim": 4, "num_heads": 1, "is_packed_batching": 1 },
339
+ "provenance": {
340
+ "source": "onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc",
341
+ "test": "ContribOpRotaryEmbeddingTest.RotaryEmbedding_CustomRotaryDim_SmallData_Phi_Packed_Batching",
342
+ "notes": "Packed batching: the sequence is longer than the cache, which upstream permits only under is_packed_batching=1. Every gathered row still indexes the cache."
343
+ },
344
+ "inputs": {
345
+ "x": {
346
+ "dtype": "float32",
347
+ "shape": [1, 3, 6],
348
+ "data": {
349
+ "kind": "values",
350
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, -1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, 1.0076, -0.7529, -0.225, -0.4327, -1.5071, -0.4586]
351
+ }
352
+ },
353
+ "positionIds": { "dtype": "uint32", "shape": [1, 3], "data": { "kind": "values", "values": [0, 0, 1] } },
354
+ "cos": { "dtype": "float32", "shape": [2, 2], "data": { "kind": "values", "values": [1.0, 1.0, 1.0, 0.5403] } },
355
+ "sin": { "dtype": "float32", "shape": [2, 2], "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.8415] } }
356
+ },
357
+ "outputs": {
358
+ "y": {
359
+ "dtype": "float32",
360
+ "shape": [1, 3, 6],
361
+ "data": {
362
+ "kind": "values",
363
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, -1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, 1.0076, -0.0427, -0.225, -0.8673, -1.5071, -0.4586]
364
+ }
365
+ }
366
+ },
367
+ "tolerance": 0.002,
368
+ "relTolerance": 0.0001
369
+ },
370
+ {
371
+ "name": "ort_rank4_interleaved_small_llama_table_f32",
372
+ "attrs": { "interleaved": 1, "num_heads": 2 },
373
+ "provenance": {
374
+ "source": "onnxruntime/test/providers/cpu/llm/rotary_embedding_op_test.cc",
375
+ "test": "RotaryEmbeddingTest.RotaryEmbedding_Interleaved_SmallData_LlamaMSFT_4D_Input",
376
+ "notes": "Rank-4 BNSH input with a rank-2 cache and a format-1 position table. The opset-23 file is the only upstream source of numerically checked rank-4 rotary data; the two ops share this math exactly."
377
+ },
378
+ "inputs": {
379
+ "x": {
380
+ "dtype": "float32",
381
+ "shape": [1, 2, 3, 4],
382
+ "data": {
383
+ "kind": "values",
384
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.2188, 1.1676, -1.0574, -0.1188, -0.811, 0.6737, -1.1233, -0.0919, -0.132, -0.2751, -0.235, 0.0937, -0.7396, -1.2425, -0.1752, 0.699, -0.6861, 0.7202, 0.1963, 0.6142]
385
+ }
386
+ },
387
+ "positionIds": { "dtype": "uint32", "shape": [1, 3], "data": { "kind": "values", "values": [0, 1, 2] } },
388
+ "cos": {
389
+ "dtype": "float32",
390
+ "shape": [8, 2],
391
+ "data": {
392
+ "kind": "values",
393
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_cos" }
394
+ }
395
+ },
396
+ "sin": {
397
+ "dtype": "float32",
398
+ "shape": [8, 2],
399
+ "data": {
400
+ "kind": "values",
401
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_small_llama_offset_f32_input_sin" }
402
+ }
403
+ }
404
+ },
405
+ "outputs": {
406
+ "y": {
407
+ "dtype": "float32",
408
+ "shape": [1, 2, 3, 4],
409
+ "data": {
410
+ "kind": "values",
411
+ "values": [-1.0408, 0.9166, -1.3042, -1.1097, -1.6411, -0.3948, -1.0561, -0.1294, -0.2751, -1.0178, -1.1212, -0.1143, -0.132, -0.2751, -0.235, 0.0937, 0.646, -1.2937, -0.1822, 0.6972, -0.3694, -0.9235, 0.184, 0.618]
412
+ }
413
+ }
414
+ },
415
+ "tolerance": 0.002,
416
+ "relTolerance": 0.0001
417
+ },
418
+ {
419
+ "name": "ort_rank4_not_interleaved_large_llama_table_f32",
420
+ "attrs": { "interleaved": 0 },
421
+ "provenance": {
422
+ "source": "onnxruntime/test/providers/cpu/llm/rotary_embedding_op_test.cc",
423
+ "test": "RotaryEmbeddingTest.RotaryEmbedding_NotInterleaved_LargeData_LlamaMSFT_4D_Input",
424
+ "notes": "Rank-4 BNSH input over two batches, four heads and eight tokens, with num_heads left unset because a rank-4 head count comes from the shape."
425
+ },
426
+ "inputs": {
427
+ "x": {
428
+ "dtype": "float32",
429
+ "shape": [2, 4, 8, 6],
430
+ "data": {
431
+ "kind": "values",
432
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_x" }
433
+ }
434
+ },
435
+ "positionIds": {
436
+ "dtype": "uint32",
437
+ "shape": [2, 8],
438
+ "data": { "kind": "values", "values": [0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7] }
439
+ },
440
+ "cos": {
441
+ "dtype": "float32",
442
+ "shape": [16, 3],
443
+ "data": {
444
+ "kind": "values",
445
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_cos" }
446
+ }
447
+ },
448
+ "sin": {
449
+ "dtype": "float32",
450
+ "shape": [16, 3],
451
+ "data": {
452
+ "kind": "values",
453
+ "values": { "$ref": "#/fixtureArrays/ort_contrib_interleaved_large_llama_offset_f32_input_sin" }
454
+ }
455
+ }
456
+ },
457
+ "outputs": {
458
+ "y": {
459
+ "dtype": "float32",
460
+ "shape": [2, 4, 8, 6],
461
+ "data": {
462
+ "kind": "values",
463
+ "values": [-1.04079998, 0.916599989, -1.30420005, -1.10969996, -1.21879995, 1.16760004, -2.10675168, 0.313278645, -1.60708773, 0.141688287, 0.0592993088, 1.58177209, -0.630788624, -0.430816084, 0.246088684, -0.174721941, 0.836671352, -0.926558971, -1.26596439, -2.2426312, -0.143917158, -1.57473576, 1.91107118, -0.659863293, 0.952363968, -0.0112946555, -0.90577817, 0.574617565, -0.583404124, -0.242907077, -1.32060885, 1.23504281, 0.760883927, 0.225809187, -0.307491541, -0.432488948, 0.0181085765, 1.12988913, -0.474278957, -0.569866478, -0.232575566, -0.647461414, 0.968330145, -0.509299397, 0.536304355, 0.91536504, 0.102920607, 0.865208685, 1.00759995, -0.752900004, -0.224999994, -0.432700008, -1.50709999, -0.458600014, -0.951666117, 0.724882483, -0.773502111, -1.74094665, 1.17627621, -0.63710475, -1.10517037, -0.524268031, -0.0400700979, -0.467021376, 0.490917653, -0.193175867, -2.36315632, -0.0442896411, -0.320379347, 1.28702021, -0.00964078028, 1.51788175, 0.51656419, 0.320925027, 0.0222803988, 0.674315631, 1.14399064, 0.886257112, 1.13239074, -0.153492898, -0.880812466, 1.86820555, -0.278367937, 0.934901953, 0.994535208, 0.827187002, 0.4941414, 1.2928561, -0.272836059, 0.929536343, 1.21685827, -0.34260717, 0.557832778, -0.992367864, 0.565743625, -1.12992167, 0.463999987, -0.498600006, 0.128900006, 2.76309991, 0.140499994, 1.11909997, 1.25286388, -0.961636603, 0.174961895, 1.70716047, -0.318457693, 0.335886538, 0.907053053, -1.02590752, 0.249643087, -0.245633602, -1.02391529, 1.87480819, -0.592516303, 1.33033943, 1.21285498, 0.13192372, -0.915585876, 1.03022671, 1.17885363, 1.77404451, -0.762661636, -1.43456602, 0.049955368, -0.27847901, 0.146011293, 2.10013723, 0.843684196, -0.153375596, 0.432110995, 0.83602643, 1.06174159, -1.55485511, -0.046079427, 0.0258956552, 0.169944048, -1.42038882, -0.0487071276, 0.315481633, 0.37001735, 0.377508819, -0.840793192, -0.63379401, -1.29439998, -0.0242999997, -0.235400006, -0.708700001, 1.1566, 0.4296, 0.154494107, -0.874685705, 0.331545562, 0.566194594, 2.07239747, 1.07093453, -0.156445935, -0.281273365, 0.711332202, 0.838858962, -0.181656986, 0.158361614, -0.79273057, -0.177007288, -0.589310288, -1.16298723, 0.0453686491, 1.18241966, 0.126825869, -0.555871427, -0.931147695, 1.45934772, -1.08596635, -0.107115202, -0.190371126, 1.33196723, -1.47011745, -0.0766584575, -0.760652125, -1.5931052, -0.0045129247, 0.70473057, 0.147792324, 0.159517035, -1.21709907, -1.65750349, 2.00992894, -0.910886765, -0.225605503, 0.452725053, -1.82546115, 0.647476315, 0.896399975, 0.571699977, -0.238999993, 0.698300004, -1.34159994, 0.271499991, 0.0294375867, 0.680094242, 0.213446647, -0.357835233, -1.6007297, 1.47927678, 0.398796827, 0.0703182518, -0.464302182, 0.485351264, -1.03685224, 0.579914272, -1.20705771, 0.0176035818, 1.53230751, 1.23830867, -0.124155864, 0.162666455, 1.44771028, -1.23949802, 1.62978542, -0.458060056, 0.660933018, 0.352941215, 1.72972739, 0.2264027, 0.729353964, -0.834230781, 0.400307, 0.0516785383, 1.61899674, -0.365789354, -2.06491113, -1.12859631, 0.320817351, -1.17251611, -0.346854568, -2.10239267, -1.61523759, 0.51734364, 0.337068677, -0.0973018557, 0.505400002, -0.668099999, -1.4382, 1.75469995, -0.960500002, -0.455799997, 0.442923486, 0.238277763, 1.15645313, -2.19831991, 1.29031682, -1.24886191, -0.509867668, 0.509775519, -3.35251856, 0.432666153, 0.741352141, -0.777529955, -2.32901859, -0.394879639, 1.3223753, 0.987909675, 0.295846343, 1.01243794, 0.505126178, 0.815001488, -1.22638965, -0.0481875092, 0.665176749, 0.90692091, 0.535472274, 1.56147265, -0.406287462, -1.7323401, -0.330429256, -1.33501053, 1.63317192, 0.490809381, 0.709373713, 0.0125124454, -0.502349257, 0.0909572691, -0.0978256166, -0.357495785, -1.35865283, 0.0379757136, -2.0119822, -0.312655121, -0.479200006, 0.675599992, -0.341300011, -0.224199995, -0.211099997, 0.628199995, -0.821949601, -1.36183679, 1.59127319, 0.725855172, -0.971916676, -0.802503109, 0.0345773101, -2.98228002, 2.41065669, 0.891961157, 0.370242298, 1.69489694, -0.107042886, 0.714565158, 2.36467719, -1.38960505, -0.699269295, -0.658058047, -0.517001033, -1.10366726, -0.720030189, 0.606771231, -0.145643681, -2.34006476, -1.26092672, -1.63743532, -0.357576013, 0.941227913, -0.569475293, 0.395344406, -1.46400166, -0.786376834, -0.865749776, 1.10432577, 0.81547302, -1.48116696, -1.24220979, 0.902649522, 0.236645028, -0.744167924, -0.482844949, 0.316913843, -1.0941, -0.368200004, -0.0163000003, -0.0644999966, -0.810100019, 0.141499996, 1.26955235, 0.626395524, -0.590327978, -0.749657393, -0.828307152, 0.73870486, 0.99614948, 0.0577319711, 0.927449882, -0.097641021, -0.0235498492, -0.216916054, 0.0532237142, 0.616131902, 0.430257797, 0.805755079, 0.485714525, 0.901634693, -0.0474238396, -0.00131507218, 1.27953064, -1.04750752, 1.23232043, 1.36800432, 0.844843626, 0.658450782, -1.20370615, -0.0611225069, -0.734763801, 0.933814406, 1.03939044, 0.0516136698, -1.31201601, -0.590313554, 0.0435673892, -0.12953417, -0.551326931, -0.740897238, -1.80020177, 0.989115238, 0.361949444, -1.45226824]
464
+ }
465
+ }
466
+ },
467
+ "tolerance": 0.002,
468
+ "relTolerance": 0.0001
469
+ },
470
+ {
471
+ "name": "quad_rank3_table_head128_heads4_batch2_f32",
472
+ "attrs": {},
473
+ "provenance": {
474
+ "notes": "Four-pair route at the Llama head size over two batches whose position rows differ, so a batch or token index mistake changes the values."
475
+ },
476
+ "inputs": {
477
+ "x": {
478
+ "dtype": "float32",
479
+ "shape": [2, 3, 512],
480
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
481
+ },
482
+ "positionIds": {
483
+ "dtype": "uint32",
484
+ "shape": [2, 3],
485
+ "data": { "kind": "values", "values": [5, 6, 7, 1, 9, 3] }
486
+ },
487
+ "cos": {
488
+ "dtype": "float32",
489
+ "shape": [16, 64],
490
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
491
+ },
492
+ "sin": {
493
+ "dtype": "float32",
494
+ "shape": [16, 64],
495
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
496
+ }
497
+ },
498
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 512] } },
499
+ "tolerance": 0.000001,
500
+ "relTolerance": 0.00001
501
+ },
502
+ {
503
+ "name": "quad_rank3_table_head128_heads4_batch2_f16",
504
+ "attrs": {},
505
+ "provenance": {
506
+ "notes": "Float16 rotary embedding with the same rank-three table geometry as the corresponding float32 case."
507
+ },
508
+ "inputs": {
509
+ "x": {
510
+ "dtype": "float16",
511
+ "shape": [2, 3, 512],
512
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
513
+ },
514
+ "positionIds": {
515
+ "dtype": "uint32",
516
+ "shape": [2, 3],
517
+ "data": { "kind": "values", "values": [5, 6, 7, 1, 9, 3] }
518
+ },
519
+ "cos": {
520
+ "dtype": "float16",
521
+ "shape": [16, 64],
522
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
523
+ },
524
+ "sin": {
525
+ "dtype": "float16",
526
+ "shape": [16, 64],
527
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
528
+ }
529
+ },
530
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 3, 512] } },
531
+ "tolerance": 0.002,
532
+ "relTolerance": 0.002
533
+ },
534
+ {
535
+ "name": "quad_rank3_table_head16_heads2_interleaved_f32",
536
+ "attrs": { "interleaved": 1 },
537
+ "provenance": {
538
+ "notes": "Interleaved pairing on the four-pair route: head size 16 is exactly two rotation groups per head, so the deinterleave and reinterleave both run."
539
+ },
540
+ "inputs": {
541
+ "x": {
542
+ "dtype": "float32",
543
+ "shape": [1, 3, 32],
544
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
545
+ },
546
+ "positionIds": { "dtype": "uint32", "shape": [1, 3], "data": { "kind": "values", "values": [2, 0, 5] } },
547
+ "cos": {
548
+ "dtype": "float32",
549
+ "shape": [8, 8],
550
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
551
+ },
552
+ "sin": {
553
+ "dtype": "float32",
554
+ "shape": [8, 8],
555
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
556
+ }
557
+ },
558
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 3, 32] } },
559
+ "tolerance": 0.000001,
560
+ "relTolerance": 0.00001
561
+ },
562
+ {
563
+ "name": "quad_rank3_offset_head16_heads2_f32",
564
+ "attrs": {},
565
+ "provenance": {
566
+ "notes": "Format-0 base offset of three: the four tokens read cache rows 3 through 6, which a token index mistake would move."
567
+ },
568
+ "inputs": {
569
+ "x": {
570
+ "dtype": "float32",
571
+ "shape": [1, 4, 32],
572
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
573
+ },
574
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [3] } },
575
+ "cos": {
576
+ "dtype": "float32",
577
+ "shape": [16, 8],
578
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
579
+ },
580
+ "sin": {
581
+ "dtype": "float32",
582
+ "shape": [16, 8],
583
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
584
+ }
585
+ },
586
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 4, 32] } },
587
+ "tolerance": 0.000001,
588
+ "relTolerance": 0.00001
589
+ },
590
+ {
591
+ "name": "quad_rank3_offset_head8_scalar_position_f32",
592
+ "attrs": {},
593
+ "provenance": { "notes": "A rank-0 scalar position id, the other spelling upstream accepts for a base offset." },
594
+ "inputs": {
595
+ "x": {
596
+ "dtype": "float32",
597
+ "shape": [1, 3, 24],
598
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
599
+ },
600
+ "positionIds": { "dtype": "uint32", "shape": [], "data": { "kind": "values", "values": [2] } },
601
+ "cos": {
602
+ "dtype": "float32",
603
+ "shape": [8, 4],
604
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
605
+ },
606
+ "sin": {
607
+ "dtype": "float32",
608
+ "shape": [8, 4],
609
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
610
+ }
611
+ },
612
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 3, 24] } },
613
+ "tolerance": 0.000001,
614
+ "relTolerance": 0.00001
615
+ },
616
+ {
617
+ "name": "quad_rank3_offset_head16_heads2_interleaved_f16",
618
+ "attrs": { "interleaved": 1 },
619
+ "provenance": { "notes": "Interleaved float16 with a base offset, the decode-loop spelling." },
620
+ "inputs": {
621
+ "x": {
622
+ "dtype": "float16",
623
+ "shape": [1, 4, 32],
624
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
625
+ },
626
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [4] } },
627
+ "cos": {
628
+ "dtype": "float16",
629
+ "shape": [16, 8],
630
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
631
+ },
632
+ "sin": {
633
+ "dtype": "float16",
634
+ "shape": [16, 8],
635
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
636
+ }
637
+ },
638
+ "outputs": { "y": { "dtype": "float16", "shape": [1, 4, 32] } },
639
+ "tolerance": 0.002,
640
+ "relTolerance": 0.002
641
+ },
642
+ {
643
+ "name": "quad_rank3_table_head80_rotary32_tail_f32",
644
+ "attrs": { "rotary_embedding_dim": 32, "num_heads": 2 },
645
+ "provenance": {
646
+ "notes": "Partial rotation at the Phi-2 geometry: 32 of 80 elements rotate and the 48-element tail is copied by the remaining rotation groups."
647
+ },
648
+ "inputs": {
649
+ "x": {
650
+ "dtype": "float32",
651
+ "shape": [1, 3, 160],
652
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
653
+ },
654
+ "positionIds": { "dtype": "uint32", "shape": [1, 3], "data": { "kind": "values", "values": [1, 4, 6] } },
655
+ "cos": {
656
+ "dtype": "float32",
657
+ "shape": [8, 16],
658
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
659
+ },
660
+ "sin": {
661
+ "dtype": "float32",
662
+ "shape": [8, 16],
663
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
664
+ }
665
+ },
666
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 3, 160] } },
667
+ "tolerance": 0.000001,
668
+ "relTolerance": 0.00001
669
+ },
670
+ {
671
+ "name": "quad_rank3_table_head24_rotary8_tail_f16",
672
+ "attrs": { "rotary_embedding_dim": 8, "num_heads": 2 },
673
+ "provenance": {
674
+ "notes": "A short rotary window inside a wide head, so most of the four-pair route's groups take the tail copy."
675
+ },
676
+ "inputs": {
677
+ "x": {
678
+ "dtype": "float16",
679
+ "shape": [2, 2, 48],
680
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
681
+ },
682
+ "positionIds": { "dtype": "uint32", "shape": [2, 2], "data": { "kind": "values", "values": [0, 3, 6, 2] } },
683
+ "cos": {
684
+ "dtype": "float16",
685
+ "shape": [8, 4],
686
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
687
+ },
688
+ "sin": {
689
+ "dtype": "float16",
690
+ "shape": [8, 4],
691
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
692
+ }
693
+ },
694
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 2, 48] } },
695
+ "tolerance": 0.002,
696
+ "relTolerance": 0.002
697
+ },
698
+ {
699
+ "name": "quad_rank3_table_head128_heads32_s8_wide_f32",
700
+ "attrs": {},
701
+ "provenance": {
702
+ "notes": "Enough rotation groups per slice that the workgroups tile the slice several times over, exercising the strided per-lane item loop rather than its guard."
703
+ },
704
+ "inputs": {
705
+ "x": {
706
+ "dtype": "float32",
707
+ "shape": [1, 8, 4096],
708
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
709
+ },
710
+ "positionIds": {
711
+ "dtype": "uint32",
712
+ "shape": [1, 8],
713
+ "data": { "kind": "values", "values": [0, 1, 2, 3, 4, 5, 6, 7] }
714
+ },
715
+ "cos": {
716
+ "dtype": "float32",
717
+ "shape": [16, 64],
718
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
719
+ },
720
+ "sin": {
721
+ "dtype": "float32",
722
+ "shape": [16, 64],
723
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
724
+ }
725
+ },
726
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 8, 4096] } },
727
+ "tolerance": 0.000001,
728
+ "relTolerance": 0.00001
729
+ },
730
+ {
731
+ "name": "quad_rank3_table_head16_rotary8_pairgate_f32",
732
+ "attrs": { "rotary_embedding_dim": 8, "num_heads": 2 },
733
+ "provenance": {
734
+ "notes": "Rotary dimension eight and the paired rotary-dimension-four case check small rotary windows at the same head width."
735
+ },
736
+ "inputs": {
737
+ "x": {
738
+ "dtype": "float32",
739
+ "shape": [1, 2, 32],
740
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
741
+ },
742
+ "positionIds": { "dtype": "uint32", "shape": [1, 2], "data": { "kind": "values", "values": [1, 5] } },
743
+ "cos": {
744
+ "dtype": "float32",
745
+ "shape": [8, 4],
746
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
747
+ },
748
+ "sin": {
749
+ "dtype": "float32",
750
+ "shape": [8, 4],
751
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
752
+ }
753
+ },
754
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 2, 32] } },
755
+ "tolerance": 0.000001,
756
+ "relTolerance": 0.00001
757
+ },
758
+ {
759
+ "name": "quad_rank4_table_head128_batch2_heads3_f32",
760
+ "attrs": {},
761
+ "provenance": {
762
+ "notes": "Rank-4 BNSH with two batches and three heads, so the pinned dispatch slice is a (batch, head) pair and every head of a batch must read the same position row."
763
+ },
764
+ "inputs": {
765
+ "x": {
766
+ "dtype": "float32",
767
+ "shape": [2, 3, 5, 128],
768
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
769
+ },
770
+ "positionIds": {
771
+ "dtype": "uint32",
772
+ "shape": [2, 5],
773
+ "data": { "kind": "values", "values": [0, 2, 4, 6, 8, 1, 3, 5, 7, 9] }
774
+ },
775
+ "cos": {
776
+ "dtype": "float32",
777
+ "shape": [16, 64],
778
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
779
+ },
780
+ "sin": {
781
+ "dtype": "float32",
782
+ "shape": [16, 64],
783
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
784
+ }
785
+ },
786
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 5, 128] } },
787
+ "tolerance": 0.000001,
788
+ "relTolerance": 0.00001
789
+ },
790
+ {
791
+ "name": "quad_rank4_table_head32_interleaved_f16",
792
+ "attrs": { "interleaved": 1 },
793
+ "provenance": { "notes": "Interleaved float16 rank-4 with non-monotonic positions." },
794
+ "inputs": {
795
+ "x": {
796
+ "dtype": "float16",
797
+ "shape": [1, 3, 4, 32],
798
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
799
+ },
800
+ "positionIds": { "dtype": "uint32", "shape": [1, 4], "data": { "kind": "values", "values": [7, 0, 3, 5] } },
801
+ "cos": {
802
+ "dtype": "float16",
803
+ "shape": [8, 16],
804
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
805
+ },
806
+ "sin": {
807
+ "dtype": "float16",
808
+ "shape": [8, 16],
809
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
810
+ }
811
+ },
812
+ "outputs": { "y": { "dtype": "float16", "shape": [1, 3, 4, 32] } },
813
+ "tolerance": 0.002,
814
+ "relTolerance": 0.002
815
+ },
816
+ {
817
+ "name": "quad_rank4_offset_head8_f32",
818
+ "attrs": {},
819
+ "provenance": {
820
+ "notes": "Rank-4 with a format-0 base offset: the token index inside the slice is the row index directly."
821
+ },
822
+ "inputs": {
823
+ "x": {
824
+ "dtype": "float32",
825
+ "shape": [1, 2, 3, 8],
826
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
827
+ },
828
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [1] } },
829
+ "cos": {
830
+ "dtype": "float32",
831
+ "shape": [8, 4],
832
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
833
+ },
834
+ "sin": {
835
+ "dtype": "float32",
836
+ "shape": [8, 4],
837
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
838
+ }
839
+ },
840
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 2, 3, 8] } },
841
+ "tolerance": 0.000001,
842
+ "relTolerance": 0.00001
843
+ },
844
+ {
845
+ "name": "quad_rank4_table_head96_rotary48_tail_f32",
846
+ "attrs": { "rotary_embedding_dim": 48, "num_heads": 2 },
847
+ "provenance": {
848
+ "notes": "The Phi-3-mini head size of 96 with a 48-element rotary window, which clears the four-pair route's rotary gate and leaves a 48-element tail."
849
+ },
850
+ "inputs": {
851
+ "x": {
852
+ "dtype": "float32",
853
+ "shape": [1, 2, 3, 96],
854
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
855
+ },
856
+ "positionIds": { "dtype": "uint32", "shape": [1, 3], "data": { "kind": "values", "values": [2, 5, 0] } },
857
+ "cos": {
858
+ "dtype": "float32",
859
+ "shape": [8, 24],
860
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
861
+ },
862
+ "sin": {
863
+ "dtype": "float32",
864
+ "shape": [8, 24],
865
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
866
+ }
867
+ },
868
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 2, 3, 96] } },
869
+ "tolerance": 0.000001,
870
+ "relTolerance": 0.00001
871
+ },
872
+ {
873
+ "name": "quad_rank4_offset_head64_batch2_f16",
874
+ "attrs": {},
875
+ "provenance": {
876
+ "notes": "Two batches sharing one base offset, so a slice-to-batch mistake shows up as a wrong head row rather than a wrong cache row."
877
+ },
878
+ "inputs": {
879
+ "x": {
880
+ "dtype": "float16",
881
+ "shape": [2, 2, 2, 64],
882
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
883
+ },
884
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [2] } },
885
+ "cos": {
886
+ "dtype": "float16",
887
+ "shape": [8, 32],
888
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
889
+ },
890
+ "sin": {
891
+ "dtype": "float16",
892
+ "shape": [8, 32],
893
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
894
+ }
895
+ },
896
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 2, 2, 64] } },
897
+ "tolerance": 0.002,
898
+ "relTolerance": 0.002
899
+ },
900
+ {
901
+ "name": "pair_rank3_table_head12_heads3_batch2_f32",
902
+ "attrs": {},
903
+ "provenance": {
904
+ "notes": "Head size 12 is not a multiple of eight, so the single-pair route carries it; two batches with different position rows."
905
+ },
906
+ "inputs": {
907
+ "x": {
908
+ "dtype": "float32",
909
+ "shape": [2, 3, 36],
910
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
911
+ },
912
+ "positionIds": {
913
+ "dtype": "uint32",
914
+ "shape": [2, 3],
915
+ "data": { "kind": "values", "values": [1, 3, 5, 8, 2, 11] }
916
+ },
917
+ "cos": {
918
+ "dtype": "float32",
919
+ "shape": [16, 6],
920
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
921
+ },
922
+ "sin": {
923
+ "dtype": "float32",
924
+ "shape": [16, 6],
925
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
926
+ }
927
+ },
928
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 36] } },
929
+ "tolerance": 0.000001,
930
+ "relTolerance": 0.00001
931
+ },
932
+ {
933
+ "name": "pair_rank3_table_head16_rotary4_pairgate_f32",
934
+ "attrs": { "rotary_embedding_dim": 4, "num_heads": 2 },
935
+ "provenance": {
936
+ "notes": "A rotary dimension of four fails the four-pair route's rotary gate even though the head size clears it, so this is the high side of that boundary."
937
+ },
938
+ "inputs": {
939
+ "x": {
940
+ "dtype": "float32",
941
+ "shape": [1, 2, 32],
942
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
943
+ },
944
+ "positionIds": { "dtype": "uint32", "shape": [1, 2], "data": { "kind": "values", "values": [1, 5] } },
945
+ "cos": {
946
+ "dtype": "float32",
947
+ "shape": [8, 2],
948
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
949
+ },
950
+ "sin": {
951
+ "dtype": "float32",
952
+ "shape": [8, 2],
953
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
954
+ }
955
+ },
956
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 2, 32] } },
957
+ "tolerance": 0.000001,
958
+ "relTolerance": 0.00001
959
+ },
960
+ {
961
+ "name": "pair_rank3_table_head7_rotary4_odd_f32",
962
+ "attrs": { "rotary_embedding_dim": 4, "num_heads": 2 },
963
+ "provenance": {
964
+ "notes": "An odd head size with an even rotary dimension, which the contrib CPU kernel supports: the last tail group copies a single element."
965
+ },
966
+ "inputs": {
967
+ "x": {
968
+ "dtype": "float32",
969
+ "shape": [1, 3, 14],
970
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
971
+ },
972
+ "positionIds": { "dtype": "uint32", "shape": [1, 3], "data": { "kind": "values", "values": [0, 4, 7] } },
973
+ "cos": {
974
+ "dtype": "float32",
975
+ "shape": [8, 2],
976
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
977
+ },
978
+ "sin": {
979
+ "dtype": "float32",
980
+ "shape": [8, 2],
981
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
982
+ }
983
+ },
984
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 3, 14] } },
985
+ "tolerance": 0.000001,
986
+ "relTolerance": 0.00001
987
+ },
988
+ {
989
+ "name": "pair_rank3_offset_head6_heads3_interleaved_f16",
990
+ "attrs": { "interleaved": 1 },
991
+ "provenance": { "notes": "Interleaved float16 on the single-pair route with a base offset." },
992
+ "inputs": {
993
+ "x": {
994
+ "dtype": "float16",
995
+ "shape": [1, 3, 18],
996
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
997
+ },
998
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [3] } },
999
+ "cos": {
1000
+ "dtype": "float16",
1001
+ "shape": [8, 3],
1002
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1003
+ },
1004
+ "sin": {
1005
+ "dtype": "float16",
1006
+ "shape": [8, 3],
1007
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1008
+ }
1009
+ },
1010
+ "outputs": { "y": { "dtype": "float16", "shape": [1, 3, 18] } },
1011
+ "tolerance": 0.002,
1012
+ "relTolerance": 0.002
1013
+ },
1014
+ {
1015
+ "name": "pair_rank4_offset_head10_f32",
1016
+ "attrs": {},
1017
+ "provenance": { "notes": "Rank-4 head size 10 on the single-pair route with a zero base offset." },
1018
+ "inputs": {
1019
+ "x": {
1020
+ "dtype": "float32",
1021
+ "shape": [1, 2, 3, 10],
1022
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
1023
+ },
1024
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [0] } },
1025
+ "cos": {
1026
+ "dtype": "float32",
1027
+ "shape": [8, 5],
1028
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1029
+ },
1030
+ "sin": {
1031
+ "dtype": "float32",
1032
+ "shape": [8, 5],
1033
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1034
+ }
1035
+ },
1036
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 2, 3, 10] } },
1037
+ "tolerance": 0.000001,
1038
+ "relTolerance": 0.00001
1039
+ },
1040
+ {
1041
+ "name": "pair_rank4_table_head9_rotary4_odd_f32",
1042
+ "attrs": { "rotary_embedding_dim": 4, "num_heads": 2 },
1043
+ "provenance": {
1044
+ "notes": "An odd rank-4 head size with a four-element rotary window; num_heads is set because upstream requires it whenever rotary_embedding_dim is, and is otherwise ignored at rank 4."
1045
+ },
1046
+ "inputs": {
1047
+ "x": {
1048
+ "dtype": "float32",
1049
+ "shape": [2, 2, 2, 9],
1050
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
1051
+ },
1052
+ "positionIds": { "dtype": "uint32", "shape": [2, 2], "data": { "kind": "values", "values": [6, 1, 0, 4] } },
1053
+ "cos": {
1054
+ "dtype": "float32",
1055
+ "shape": [8, 2],
1056
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1057
+ },
1058
+ "sin": {
1059
+ "dtype": "float32",
1060
+ "shape": [8, 2],
1061
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1062
+ }
1063
+ },
1064
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 2, 9] } },
1065
+ "tolerance": 0.000001,
1066
+ "relTolerance": 0.00001
1067
+ },
1068
+ {
1069
+ "name": "pair_rank3_table_head6_heads5_hidden30_f32",
1070
+ "attrs": {},
1071
+ "provenance": {
1072
+ "notes": "Five heads of six with num_heads omitted, so the head count is inferred as hidden_size divided by twice the cache width."
1073
+ },
1074
+ "inputs": {
1075
+ "x": {
1076
+ "dtype": "float32",
1077
+ "shape": [1, 4, 30],
1078
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
1079
+ },
1080
+ "positionIds": { "dtype": "uint32", "shape": [1, 4], "data": { "kind": "values", "values": [0, 1, 2, 3] } },
1081
+ "cos": {
1082
+ "dtype": "float32",
1083
+ "shape": [8, 3],
1084
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1085
+ },
1086
+ "sin": {
1087
+ "dtype": "float32",
1088
+ "shape": [8, 3],
1089
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1090
+ }
1091
+ },
1092
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 4, 30] } },
1093
+ "tolerance": 0.000001,
1094
+ "relTolerance": 0.00001
1095
+ },
1096
+ {
1097
+ "name": "decode_rank3_offset_one_token_f32",
1098
+ "attrs": {},
1099
+ "provenance": {
1100
+ "notes": "Single-token decode with a one-element position vector, the shape a decode loop emits."
1101
+ },
1102
+ "inputs": {
1103
+ "x": {
1104
+ "dtype": "float32",
1105
+ "shape": [1, 1, 4096],
1106
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
1107
+ },
1108
+ "positionIds": { "dtype": "uint32", "shape": [1], "data": { "kind": "values", "values": [17] } },
1109
+ "cos": {
1110
+ "dtype": "float32",
1111
+ "shape": [32, 64],
1112
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1113
+ },
1114
+ "sin": {
1115
+ "dtype": "float32",
1116
+ "shape": [32, 64],
1117
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1118
+ }
1119
+ },
1120
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 1, 4096] } },
1121
+ "tolerance": 0.000001,
1122
+ "relTolerance": 0.00001
1123
+ },
1124
+ {
1125
+ "name": "decode_rank3_table_one_token_f32",
1126
+ "attrs": {},
1127
+ "provenance": {
1128
+ "notes": "The same single token spelled as a 1x1 position table. Upstream reads a rank-2 tensor as the table format however few elements it holds, so this is the other side of the format boundary."
1129
+ },
1130
+ "inputs": {
1131
+ "x": {
1132
+ "dtype": "float32",
1133
+ "shape": [1, 1, 4096],
1134
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
1135
+ },
1136
+ "positionIds": { "dtype": "uint32", "shape": [1, 1], "data": { "kind": "values", "values": [17] } },
1137
+ "cos": {
1138
+ "dtype": "float32",
1139
+ "shape": [32, 64],
1140
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1141
+ },
1142
+ "sin": {
1143
+ "dtype": "float32",
1144
+ "shape": [32, 64],
1145
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1146
+ }
1147
+ },
1148
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 1, 4096] } },
1149
+ "tolerance": 0.000001,
1150
+ "relTolerance": 0.00001
1151
+ },
1152
+ {
1153
+ "name": "decode_rank3_table_kv_heads8_f32",
1154
+ "attrs": {},
1155
+ "provenance": {
1156
+ "notes": "The grouped-query key projection of the same decode step: eight heads instead of thirty-two."
1157
+ },
1158
+ "inputs": {
1159
+ "x": {
1160
+ "dtype": "float32",
1161
+ "shape": [1, 1, 1024],
1162
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 1.0 }
1163
+ },
1164
+ "positionIds": { "dtype": "uint32", "shape": [1, 1], "data": { "kind": "values", "values": [21] } },
1165
+ "cos": {
1166
+ "dtype": "float32",
1167
+ "shape": [32, 64],
1168
+ "data": { "kind": "rotaryCos", "thetaStart": 0.1, "thetaStep": 0.2 }
1169
+ },
1170
+ "sin": {
1171
+ "dtype": "float32",
1172
+ "shape": [32, 64],
1173
+ "data": { "kind": "rotarySin", "thetaStart": 0.1, "thetaStep": 0.2 }
1174
+ }
1175
+ },
1176
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 1, 1024] } },
1177
+ "tolerance": 0.000001,
1178
+ "relTolerance": 0.00001
1179
+ }
1180
+ ]
1181
+ }