Xenova HF Staff commited on
Commit
8929c9c
·
verified ·
1 Parent(s): 296e50f

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -12,7 +12,7 @@ tags:
12
 
13
  ## Description
14
 
15
- Multimodal rotary position embedding (M-RoPE) for Qwen models. Each token has temporal, height, and width position streams; `mrope_section` partitions the half-rotary axis and `mrope_layout` assigns them. Text-only tokens set all streams equal, reducing the op to `RotaryEmbedding`. The effective rotary dimension must be positive and even; an odd head size is supported with a smaller even `rotary_embedding_dim`. This package supports float16/float32 and non-packed mode; bfloat16 and packed batching are not implemented. Position ids must be valid non-negative cache-row indices.
16
 
17
  See the [ONNX Runtime `MRotaryEmbedding` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.MRotaryEmbedding) for the reference semantics.
18
 
@@ -21,7 +21,7 @@ See the [ONNX Runtime `MRotaryEmbedding` contrib-operator spec](https://github.c
21
  | Name | Upstream name | Logical dtype | WebGPU storage | Rank | Shape | Description | Presence |
22
  | --- | --- | --- | --- | --- | --- | --- | --- |
23
  | `x` | `input` | `T` | same as logical dtype | — | — | Input token embeddings. Shape is `(batch_size, sequence_length, hidden_size)` for rank 3 or `(batch_size, num_heads, sequence_length, head_size)` for rank 4. The effective rotary dimension must be even, and `num_heads` is required for rank-3 input. | required |
24
- | `positionIds` | `position_ids` | `M` | `uint32` | `3` | — | Logical int64 position indices of shape `(3, batch_size, sequence_length)`, holding the temporal, height and width streams in that order along the first axis. Valid positions are non-negative cache-row indices and use uint32 WebGPU storage. | required |
25
  | `cos` | `cos_cache` | `T` | same as logical dtype | `2` | — | Precomputed cosine values of shape `(max_sequence_length, rotary_dim/2)`, shared by all three position streams. | 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
 
@@ -52,18 +52,25 @@ Attributes and default values (overridable per request):
52
  | `T` | `float32`, `float16` |
53
  | `M` | `int64` |
54
 
 
 
 
 
 
 
55
  ## Files
56
 
57
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
58
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
59
  - [`test.json`](build/webgpu/test.json) — correctness cases
60
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
 
61
  - [`mrotary-embedding.wgsl.jinja`](build/webgpu/mrotary-embedding.wgsl.jinja)
62
 
63
  ## Use with `@huggingface/kernels`
64
 
65
  ```sh
66
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
67
  ```
68
 
69
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
@@ -78,10 +85,10 @@ import { getKernel } from "@huggingface/kernels";
78
 
79
  const kernel = await getKernel("webgpu-kernels/com.microsoft.MRotaryEmbedding", { version: 1 });
80
  const { y } = await kernel({
81
- x: { data: xData, shape: [1, 4, 16] },
82
- positionIds: { data: positionIdsData, shape: [3, 1, 4] },
83
- cos: { data: cosData, shape: [8, 4] },
84
- sin: { data: sinData, shape: [8, 4] },
85
  }, {
86
  attrs: { num_heads: 2, mrope_section: [2, 1, 1] },
87
  });
 
12
 
13
  ## Description
14
 
15
+ Multimodal rotary position embedding (M-RoPE) for Qwen models. Each token has temporal, height, and width position streams; `mrope_section` partitions the half-rotary axis and `mrope_layout` assigns them. Text-only tokens set all streams equal, reducing the op to `RotaryEmbedding`. The effective rotary dimension must be positive and even; an odd head size is supported with a smaller even `rotary_embedding_dim`. This package supports float16/float32 and non-packed mode; bfloat16 and packed batching are not implemented. An out-of-range position id copies its rotation pair through unchanged.
16
 
17
  See the [ONNX Runtime `MRotaryEmbedding` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.MRotaryEmbedding) for the reference semantics.
18
 
 
21
  | Name | Upstream name | Logical dtype | WebGPU storage | Rank | Shape | Description | Presence |
22
  | --- | --- | --- | --- | --- | --- | --- | --- |
23
  | `x` | `input` | `T` | same as logical dtype | — | — | Input token embeddings. Shape is `(batch_size, sequence_length, hidden_size)` for rank 3 or `(batch_size, num_heads, sequence_length, head_size)` for rank 4. The effective rotary dimension must be even, and `num_heads` is required for rank-3 input. | required |
24
+ | `positionIds` | `position_ids` | `M` | `int32` | `3` | — | Logical int64 position indices of shape `(3, batch_size, sequence_length)` for temporal, height and width streams. Signed int32 storage saturates larger values; negative or out-of-cache indices copy the corresponding pair unchanged. | required |
25
  | `cos` | `cos_cache` | `T` | same as logical dtype | `2` | — | Precomputed cosine values of shape `(max_sequence_length, rotary_dim/2)`, shared by all three position streams. | 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
 
 
52
  | `T` | `float32`, `float16` |
53
  | `M` | `int64` |
54
 
55
+ ## Implementation variants
56
+
57
+ One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
58
+
59
+ - `pairs4` — Four contiguous split-half pairs share vector activation loads and stores. Each component gathers its own position stream and preserves scaled-coefficient rounding to the storage type.
60
+
61
  ## Files
62
 
63
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
64
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
65
  - [`test.json`](build/webgpu/test.json) — correctness cases
66
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
67
+ - [`mrotary-embedding-vector.wgsl.jinja`](build/webgpu/mrotary-embedding-vector.wgsl.jinja)
68
  - [`mrotary-embedding.wgsl.jinja`](build/webgpu/mrotary-embedding.wgsl.jinja)
69
 
70
  ## Use with `@huggingface/kernels`
71
 
72
  ```sh
73
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
74
  ```
75
 
76
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
85
 
86
  const kernel = await getKernel("webgpu-kernels/com.microsoft.MRotaryEmbedding", { version: 1 });
87
  const { y } = await kernel({
88
+ x: { data: xData, shape: [2, 2, 16] },
89
+ positionIds: { data: positionIdsData, shape: [3, 2, 2] },
90
+ cos: { data: cosData, shape: [2, 4] },
91
+ sin: { data: sinData, shape: [2, 4] },
92
  }, {
93
  attrs: { num_heads: 2, mrope_section: [2, 1, 1] },
94
  });
build/webgpu/bench.json CHANGED
@@ -14,7 +14,7 @@
14
  "vars": { "elems": 16384, "rotPairs": 8192 },
15
  "inputs": {
16
  "x": { "shape": [1, 4, 64, 64], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
17
- "positionIds": { "shape": [3, 1, 64], "dtype": "uint32", "dist": "uniform", "seed": 1001, "min": 0, "max": 127 },
18
  "cos": { "shape": [128, 32], "dtype": "float32", "dist": "uniform", "seed": 1002, "scale": 0.1, "offset": 0.9 },
19
  "sin": { "shape": [128, 32], "dtype": "float32", "dist": "uniform", "seed": 1003, "scale": 0.1 }
20
  },
@@ -41,14 +41,7 @@
41
  },
42
  "inputs": {
43
  "x": { "shape": [1, 16, 1, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
44
- "positionIds": {
45
- "shape": [3, 1, 1],
46
- "dtype": "uint32",
47
- "dist": "uniform",
48
- "seed": 1001,
49
- "min": 0,
50
- "max": 32767
51
- },
52
  "cos": {
53
  "shape": [32768, 64],
54
  "dtype": "float32",
@@ -84,7 +77,7 @@
84
  "x": { "shape": [1, 16, 512, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
85
  "positionIds": {
86
  "shape": [3, 1, 512],
87
- "dtype": "uint32",
88
  "dist": "uniform",
89
  "seed": 1001,
90
  "min": 0,
@@ -123,14 +116,7 @@
123
  },
124
  "inputs": {
125
  "x": { "shape": [1, 8, 1, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
126
- "positionIds": {
127
- "shape": [3, 1, 1],
128
- "dtype": "uint32",
129
- "dist": "uniform",
130
- "seed": 1001,
131
- "min": 0,
132
- "max": 32767
133
- },
134
  "cos": {
135
  "shape": [32768, 64],
136
  "dtype": "float32",
@@ -166,7 +152,7 @@
166
  "x": { "shape": [1, 16, 512, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
167
  "positionIds": {
168
  "shape": [3, 1, 512],
169
- "dtype": "uint32",
170
  "dist": "uniform",
171
  "seed": 1001,
172
  "min": 0,
@@ -207,7 +193,7 @@
207
  "x": { "shape": [1, 16, 512, 128], "dtype": "float16", "dist": "normal", "seed": 1000, "scale": 0.2 },
208
  "positionIds": {
209
  "shape": [3, 1, 512],
210
- "dtype": "uint32",
211
  "dist": "uniform",
212
  "seed": 1001,
213
  "min": 0,
@@ -244,7 +230,7 @@
244
  "x": { "shape": [4, 32, 2048, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
245
  "positionIds": {
246
  "shape": [3, 4, 2048],
247
- "dtype": "uint32",
248
  "dist": "uniform",
249
  "seed": 1001,
250
  "min": 0,
@@ -265,6 +251,114 @@
265
  "primary": false,
266
  "metrics": [{ "type": "bandwidth", "value": "args.elems * 4 * 2 + args.rotPairs * 4 * 2" }]
267
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
268
  }
269
  ]
270
  }
 
14
  "vars": { "elems": 16384, "rotPairs": 8192 },
15
  "inputs": {
16
  "x": { "shape": [1, 4, 64, 64], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
17
+ "positionIds": { "shape": [3, 1, 64], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 127 },
18
  "cos": { "shape": [128, 32], "dtype": "float32", "dist": "uniform", "seed": 1002, "scale": 0.1, "offset": 0.9 },
19
  "sin": { "shape": [128, 32], "dtype": "float32", "dist": "uniform", "seed": 1003, "scale": 0.1 }
20
  },
 
41
  },
42
  "inputs": {
43
  "x": { "shape": [1, 16, 1, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
44
+ "positionIds": { "shape": [3, 1, 1], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 32767 },
 
 
 
 
 
 
 
45
  "cos": {
46
  "shape": [32768, 64],
47
  "dtype": "float32",
 
77
  "x": { "shape": [1, 16, 512, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
78
  "positionIds": {
79
  "shape": [3, 1, 512],
80
+ "dtype": "int32",
81
  "dist": "uniform",
82
  "seed": 1001,
83
  "min": 0,
 
116
  },
117
  "inputs": {
118
  "x": { "shape": [1, 8, 1, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
119
+ "positionIds": { "shape": [3, 1, 1], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 32767 },
 
 
 
 
 
 
 
120
  "cos": {
121
  "shape": [32768, 64],
122
  "dtype": "float32",
 
152
  "x": { "shape": [1, 16, 512, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
153
  "positionIds": {
154
  "shape": [3, 1, 512],
155
+ "dtype": "int32",
156
  "dist": "uniform",
157
  "seed": 1001,
158
  "min": 0,
 
193
  "x": { "shape": [1, 16, 512, 128], "dtype": "float16", "dist": "normal", "seed": 1000, "scale": 0.2 },
194
  "positionIds": {
195
  "shape": [3, 1, 512],
196
+ "dtype": "int32",
197
  "dist": "uniform",
198
  "seed": 1001,
199
  "min": 0,
 
230
  "x": { "shape": [4, 32, 2048, 128], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
231
  "positionIds": {
232
  "shape": [3, 4, 2048],
233
+ "dtype": "int32",
234
  "dist": "uniform",
235
  "seed": 1001,
236
  "min": 0,
 
251
  "primary": false,
252
  "metrics": [{ "type": "bandwidth", "value": "args.elems * 4 * 2 + args.rotPairs * 4 * 2" }]
253
  }
254
+ },
255
+ {
256
+ "name": "mrope-vector-neighbor-rank3-d136-float32",
257
+ "preset": "model",
258
+ "attrs": {
259
+ "mrope_layout": 0,
260
+ "interleaved": 0,
261
+ "mrope_section": [1, 7, 60],
262
+ "rotary_embedding_dim": 136,
263
+ "scale": 0.5,
264
+ "num_heads": 5
265
+ },
266
+ "vars": { "elems": 349520, "rotPairs": 174760, "dtype": "float32" },
267
+ "provenance": {
268
+ "notes": "Aligned non-power-of-two head width, multiple batches, odd sequence, section-crossing vectors and non-unit cache scale. Counts logical input/output plus coefficient loads, excluding position indices and cache reuse."
269
+ },
270
+ "inputs": {
271
+ "x": { "shape": [2, 257, 680], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
272
+ "positionIds": { "shape": [3, 2, 257], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 512 },
273
+ "cos": { "shape": [513, 68], "dtype": "float32", "dist": "uniform", "seed": 1002, "scale": 0.1, "offset": 0.9 },
274
+ "sin": { "shape": [513, 68], "dtype": "float32", "dist": "uniform", "seed": 1003, "scale": 0.1 }
275
+ },
276
+ "outputs": { "y": { "shape": [2, 257, 680], "dtype": "float32" } },
277
+ "bench": {
278
+ "primary": false,
279
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.rotPairs * 2) * dtypeBytes(args.dtype)" }]
280
+ }
281
+ },
282
+ {
283
+ "name": "mrope-vector-neighbor-rank4-d136-float32",
284
+ "preset": "model",
285
+ "attrs": {
286
+ "mrope_layout": 1,
287
+ "interleaved": 0,
288
+ "mrope_section": [1, 7, 60],
289
+ "rotary_embedding_dim": 136,
290
+ "scale": 0.5,
291
+ "num_heads": 5
292
+ },
293
+ "vars": { "elems": 349520, "rotPairs": 174760, "dtype": "float32" },
294
+ "provenance": {
295
+ "notes": "Aligned non-power-of-two head width, multiple batches, odd sequence, section-crossing vectors and non-unit cache scale. Counts logical input/output plus coefficient loads, excluding position indices and cache reuse."
296
+ },
297
+ "inputs": {
298
+ "x": { "shape": [2, 5, 257, 136], "dtype": "float32", "dist": "normal", "seed": 1000, "scale": 0.2 },
299
+ "positionIds": { "shape": [3, 2, 257], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 512 },
300
+ "cos": { "shape": [513, 68], "dtype": "float32", "dist": "uniform", "seed": 1002, "scale": 0.1, "offset": 0.9 },
301
+ "sin": { "shape": [513, 68], "dtype": "float32", "dist": "uniform", "seed": 1003, "scale": 0.1 }
302
+ },
303
+ "outputs": { "y": { "shape": [2, 5, 257, 136], "dtype": "float32" } },
304
+ "bench": {
305
+ "primary": false,
306
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.rotPairs * 2) * dtypeBytes(args.dtype)" }]
307
+ }
308
+ },
309
+ {
310
+ "name": "mrope-vector-neighbor-rank3-d136-float16",
311
+ "preset": "model",
312
+ "attrs": {
313
+ "mrope_layout": 0,
314
+ "interleaved": 0,
315
+ "mrope_section": [1, 7, 60],
316
+ "rotary_embedding_dim": 136,
317
+ "scale": 0.5,
318
+ "num_heads": 5
319
+ },
320
+ "vars": { "elems": 349520, "rotPairs": 174760, "dtype": "float16" },
321
+ "provenance": {
322
+ "notes": "Aligned non-power-of-two head width, multiple batches, odd sequence, section-crossing vectors and non-unit cache scale. Counts logical input/output plus coefficient loads, excluding position indices and cache reuse."
323
+ },
324
+ "inputs": {
325
+ "x": { "shape": [2, 257, 680], "dtype": "float16", "dist": "normal", "seed": 1000, "scale": 0.2 },
326
+ "positionIds": { "shape": [3, 2, 257], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 512 },
327
+ "cos": { "shape": [513, 68], "dtype": "float16", "dist": "uniform", "seed": 1002, "scale": 0.1, "offset": 0.9 },
328
+ "sin": { "shape": [513, 68], "dtype": "float16", "dist": "uniform", "seed": 1003, "scale": 0.1 }
329
+ },
330
+ "outputs": { "y": { "shape": [2, 257, 680], "dtype": "float16" } },
331
+ "bench": {
332
+ "primary": false,
333
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.rotPairs * 2) * dtypeBytes(args.dtype)" }]
334
+ }
335
+ },
336
+ {
337
+ "name": "mrope-vector-neighbor-rank4-d136-float16",
338
+ "preset": "model",
339
+ "attrs": {
340
+ "mrope_layout": 1,
341
+ "interleaved": 0,
342
+ "mrope_section": [1, 7, 60],
343
+ "rotary_embedding_dim": 136,
344
+ "scale": 0.5,
345
+ "num_heads": 5
346
+ },
347
+ "vars": { "elems": 349520, "rotPairs": 174760, "dtype": "float16" },
348
+ "provenance": {
349
+ "notes": "Aligned non-power-of-two head width, multiple batches, odd sequence, section-crossing vectors and non-unit cache scale. Counts logical input/output plus coefficient loads, excluding position indices and cache reuse."
350
+ },
351
+ "inputs": {
352
+ "x": { "shape": [2, 5, 257, 136], "dtype": "float16", "dist": "normal", "seed": 1000, "scale": 0.2 },
353
+ "positionIds": { "shape": [3, 2, 257], "dtype": "int32", "dist": "uniform", "seed": 1001, "min": 0, "max": 512 },
354
+ "cos": { "shape": [513, 68], "dtype": "float16", "dist": "uniform", "seed": 1002, "scale": 0.1, "offset": 0.9 },
355
+ "sin": { "shape": [513, 68], "dtype": "float16", "dist": "uniform", "seed": 1003, "scale": 0.1 }
356
+ },
357
+ "outputs": { "y": { "shape": [2, 5, 257, 136], "dtype": "float16" } },
358
+ "bench": {
359
+ "primary": false,
360
+ "metrics": [{ "type": "bandwidth", "value": "(args.elems * 2 + args.rotPairs * 2) * dtypeBytes(args.dtype)" }]
361
+ }
362
  }
363
  ]
364
  }
build/webgpu/manifest.json CHANGED
@@ -4,7 +4,7 @@
4
  "sinceVersion": 1,
5
  "inputs": {
6
  "x": { "onnx": "input", "dtype": "T" },
7
- "positionIds": { "onnx": "position_ids", "dtype": "M", "rank": 3, "storage": "uint32", "narrowing": "checked" },
8
  "cos": { "onnx": "cos_cache", "dtype": "T", "rank": 2 },
9
  "sin": { "onnx": "sin_cache", "dtype": "T", "rank": 2 }
10
  },
@@ -37,91 +37,91 @@
37
  "sectionsDefined": "attrs.mrope_section is defined and (attrs.mrope_section | length) == 3",
38
  "sectionsValid": "sectionsDefined and attrs.mrope_section[0] >= 0 and attrs.mrope_section[1] >= 0 and attrs.mrope_section[2] >= 0 and attrs.mrope_section[0] + attrs.mrope_section[1] + attrs.mrope_section[2] == halfRotaryDim",
39
  "commonContract": "f16Ok(dtypes.T) and sameShape(shapes.x, shapes.y) and (attrs.interleaved == 0 or attrs.interleaved == 1) and (attrs.mrope_layout == 0 or attrs.mrope_layout == 1) and attrs.num_heads >= 0 and attrs.rotary_embedding_dim >= 0 and (attrs.rotary_embedding_dim == 0 or attrs.num_heads > 0) and sameShape(shapes.cos, shapes.sin) and ranks.cos == 2 and ranks.sin == 2 and sectionsValid and ranks.positionIds == 3 and dim(shapes.positionIds, 0) == 3 and dim(shapes.positionIds, 1) == dim(shapes.x, 0) and dim(shapes.positionIds, 2) == dim(shapes.x, 1 if ranks.x == 3 else 2)",
40
- "rank3Contract": "commonContract and ranks.x == 3 and ranks.y == 3 and attrs.num_heads is defined and attrs.num_heads >= 1 and dim(shapes.x, 2) % attrs.num_heads == 0 and rank3HeadSize > 0 and effectiveRotaryDim % 2 == 0 and (attrs.rotary_embedding_dim == 0 or attrs.rotary_embedding_dim <= rank3HeadSize) and halfRotaryDim * 2 == effectiveRotaryDim",
41
  "rank4Contract": "commonContract and ranks.x == 4 and ranks.y == 4 and dim(shapes.x, 3) > 0 and effectiveRotaryDim % 2 == 0 and (attrs.rotary_embedding_dim == 0 or attrs.rotary_embedding_dim <= dim(shapes.x, 3)) and halfRotaryDim * 2 == effectiveRotaryDim"
42
  },
43
- "when": ["pairDispatchOk"],
44
  "bindings": {
45
- "x": { "buffer": "read-only-storage", "elementType": "$T" },
46
- "position_ids": { "arg": "positionIds", "buffer": "read-only-storage", "elementType": "u32" },
47
- "cos_cache": { "arg": "cos", "buffer": "read-only-storage", "elementType": "$T" },
48
- "sin_cache": { "arg": "sin", "buffer": "read-only-storage", "elementType": "$T" },
49
- "y": { "buffer": "storage", "elementType": "$T" },
50
  "params": {
51
- "buffer": "uniform",
52
  "struct": [
53
  { "name": "pairCount", "type": "u32", "value": "pairCount" },
54
  { "name": "batchSize", "type": "u32", "value": "dim(shapes.x, 0)" },
55
- { "name": "sequenceLength", "type": "u32", "value": "dim(shapes.x, 1)" },
56
- { "name": "numHeads", "type": "u32", "value": "attrs.num_heads" },
57
- { "name": "headSize", "type": "u32", "value": "rank3HeadSize" },
58
  { "name": "rotaryDim", "type": "u32", "value": "halfRotaryDim * 2" },
59
  { "name": "halfRotaryDim", "type": "u32", "value": "halfRotaryDim" },
60
  { "name": "section0", "type": "u32", "value": "attrs.mrope_section[0]" },
61
  { "name": "section1", "type": "u32", "value": "attrs.mrope_section[1]" },
62
  { "name": "section2", "type": "u32", "value": "attrs.mrope_section[2]" },
63
- { "name": "scale", "type": "f32", "value": "attrs.scale" }
 
64
  ]
65
  },
66
- "params_2": {
67
- "name": "params",
68
- "buffer": "uniform",
69
  "struct": [
70
  { "name": "pairCount", "type": "u32", "value": "pairCount" },
71
  { "name": "batchSize", "type": "u32", "value": "dim(shapes.x, 0)" },
72
- { "name": "sequenceLength", "type": "u32", "value": "dim(shapes.x, 2)" },
73
- { "name": "numHeads", "type": "u32", "value": "dim(shapes.x, 1)" },
74
- { "name": "headSize", "type": "u32", "value": "dim(shapes.x, 3)" },
75
- { "name": "rotaryDim", "type": "u32", "value": "halfRotaryDim * 2" },
76
  { "name": "halfRotaryDim", "type": "u32", "value": "halfRotaryDim" },
77
  { "name": "section0", "type": "u32", "value": "attrs.mrope_section[0]" },
78
  { "name": "section1", "type": "u32", "value": "attrs.mrope_section[1]" },
79
  { "name": "section2", "type": "u32", "value": "attrs.mrope_section[2]" },
80
- { "name": "scale", "type": "f32", "value": "attrs.scale" }
81
- ]
 
 
82
  }
83
  },
84
  "variants": [
85
  {
86
- "id": "rank3",
87
- "when": ["rank3Contract"],
88
  "derive": {
89
- "interleaved": "attrs.interleaved != 0",
90
  "mropeSectioned": "attrs.mrope_layout == 0",
91
- "usesF16": "dtypes.T == \"f16\"",
 
92
  "scalar": "dtypes.T"
93
  },
94
  "passes": [
95
  {
96
  "id": "main",
97
- "name": "mrotary_embedding3d",
98
- "shader": "mrotary-embedding.wgsl.jinja",
99
- "derive": { "rank": 3 },
100
- "bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params"],
101
  "dispatch": {
102
- "x": "min(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)",
103
- "y": "ceilDiv(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)",
104
  "z": 1
105
  }
106
  }
107
- ]
 
108
  },
109
  {
110
- "id": "rank4",
111
- "when": ["rank4Contract"],
112
  "derive": {
113
  "interleaved": "attrs.interleaved != 0",
114
  "mropeSectioned": "attrs.mrope_layout == 0",
115
- "usesF16": "dtypes.T == \"f16\"",
116
  "scalar": "dtypes.T"
117
  },
118
  "passes": [
119
  {
120
  "id": "main",
121
- "name": "mrotary_embedding4d",
122
  "shader": "mrotary-embedding.wgsl.jinja",
123
- "derive": { "rank": 4 },
124
- "bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params_2"],
125
  "dispatch": {
126
  "x": "min(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)",
127
  "y": "ceilDiv(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)",
 
4
  "sinceVersion": 1,
5
  "inputs": {
6
  "x": { "onnx": "input", "dtype": "T" },
7
+ "positionIds": { "onnx": "position_ids", "dtype": "M", "rank": 3, "storage": "int32", "narrowing": "saturating" },
8
  "cos": { "onnx": "cos_cache", "dtype": "T", "rank": 2 },
9
  "sin": { "onnx": "sin_cache", "dtype": "T", "rank": 2 }
10
  },
 
37
  "sectionsDefined": "attrs.mrope_section is defined and (attrs.mrope_section | length) == 3",
38
  "sectionsValid": "sectionsDefined and attrs.mrope_section[0] >= 0 and attrs.mrope_section[1] >= 0 and attrs.mrope_section[2] >= 0 and attrs.mrope_section[0] + attrs.mrope_section[1] + attrs.mrope_section[2] == halfRotaryDim",
39
  "commonContract": "f16Ok(dtypes.T) and sameShape(shapes.x, shapes.y) and (attrs.interleaved == 0 or attrs.interleaved == 1) and (attrs.mrope_layout == 0 or attrs.mrope_layout == 1) and attrs.num_heads >= 0 and attrs.rotary_embedding_dim >= 0 and (attrs.rotary_embedding_dim == 0 or attrs.num_heads > 0) and sameShape(shapes.cos, shapes.sin) and ranks.cos == 2 and ranks.sin == 2 and sectionsValid and ranks.positionIds == 3 and dim(shapes.positionIds, 0) == 3 and dim(shapes.positionIds, 1) == dim(shapes.x, 0) and dim(shapes.positionIds, 2) == dim(shapes.x, 1 if ranks.x == 3 else 2)",
40
+ "rank3Contract": "commonContract and ranks.x == 3 and ranks.y == 3 and attrs.num_heads is defined and attrs.num_heads >= 1 and dim(shapes.x, 2) % attrs.num_heads == 0 and (rank3HeadSize > 0 or numel(shapes.x) == 0) and effectiveRotaryDim % 2 == 0 and (attrs.rotary_embedding_dim == 0 or attrs.rotary_embedding_dim <= rank3HeadSize) and halfRotaryDim * 2 == effectiveRotaryDim",
41
  "rank4Contract": "commonContract and ranks.x == 4 and ranks.y == 4 and dim(shapes.x, 3) > 0 and effectiveRotaryDim % 2 == 0 and (attrs.rotary_embedding_dim == 0 or attrs.rotary_embedding_dim <= dim(shapes.x, 3)) and halfRotaryDim * 2 == effectiveRotaryDim"
42
  },
43
+ "when": ["pairDispatchOk", "rank3Contract or rank4Contract"],
44
  "bindings": {
45
+ "x": { "elementType": "$T" },
46
+ "position_ids": { "arg": "positionIds", "elementType": "i32" },
47
+ "cos_cache": { "arg": "cos", "elementType": "$T" },
48
+ "sin_cache": { "arg": "sin", "elementType": "$T" },
49
+ "y": { "elementType": "$T" },
50
  "params": {
 
51
  "struct": [
52
  { "name": "pairCount", "type": "u32", "value": "pairCount" },
53
  { "name": "batchSize", "type": "u32", "value": "dim(shapes.x, 0)" },
54
+ { "name": "sequenceLength", "type": "u32", "value": "dim(shapes.x, 1 if ranks.x == 3 else 2)" },
55
+ { "name": "numHeads", "type": "u32", "value": "attrs.num_heads if ranks.x == 3 else dim(shapes.x, 1)" },
56
+ { "name": "headSize", "type": "u32", "value": "headSize" },
57
  { "name": "rotaryDim", "type": "u32", "value": "halfRotaryDim * 2" },
58
  { "name": "halfRotaryDim", "type": "u32", "value": "halfRotaryDim" },
59
  { "name": "section0", "type": "u32", "value": "attrs.mrope_section[0]" },
60
  { "name": "section1", "type": "u32", "value": "attrs.mrope_section[1]" },
61
  { "name": "section2", "type": "u32", "value": "attrs.mrope_section[2]" },
62
+ { "name": "scale", "type": "f32", "value": "attrs.scale" },
63
+ { "name": "maxSequenceLength", "type": "u32", "value": "dim(shapes.cos, 0)" }
64
  ]
65
  },
66
+ "x_vector": { "arg": "x", "name": "x", "elementType": "$vectorScalar" },
67
+ "y_vector": { "arg": "y", "name": "y", "elementType": "$vectorScalar" },
68
+ "params_vector": {
69
  "struct": [
70
  { "name": "pairCount", "type": "u32", "value": "pairCount" },
71
  { "name": "batchSize", "type": "u32", "value": "dim(shapes.x, 0)" },
72
+ { "name": "sequenceLength", "type": "u32", "value": "dim(shapes.x, 1 if ranks.x == 3 else 2)" },
73
+ { "name": "numHeads", "type": "u32", "value": "attrs.num_heads if ranks.x == 3 else dim(shapes.x, 1)" },
74
+ { "name": "headSize", "type": "u32", "value": "headSize" },
 
75
  { "name": "halfRotaryDim", "type": "u32", "value": "halfRotaryDim" },
76
  { "name": "section0", "type": "u32", "value": "attrs.mrope_section[0]" },
77
  { "name": "section1", "type": "u32", "value": "attrs.mrope_section[1]" },
78
  { "name": "section2", "type": "u32", "value": "attrs.mrope_section[2]" },
79
+ { "name": "scale", "type": "f32", "value": "attrs.scale" },
80
+ { "name": "maxSequenceLength", "type": "u32", "value": "dim(shapes.cos, 0)" }
81
+ ],
82
+ "name": "params"
83
  }
84
  },
85
  "variants": [
86
  {
87
+ "id": "pairs4",
88
+ "when": ["headSize % 8 == 0", "effectiveRotaryDim == headSize", "attrs.interleaved == 0"],
89
  "derive": {
 
90
  "mropeSectioned": "attrs.mrope_layout == 0",
91
+ "pairWidth": "4",
92
+ "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
93
  "scalar": "dtypes.T"
94
  },
95
  "passes": [
96
  {
97
  "id": "main",
98
+ "name": "MRotaryEmbedding.Vector",
99
+ "shader": "mrotary-embedding-vector.wgsl.jinja",
100
+ "derive": { "rank": "ranks.x" },
101
+ "bindings": ["x_vector", "position_ids", "cos_cache", "sin_cache", "y_vector", "params_vector"],
102
  "dispatch": {
103
+ "x": "min(ceilDiv((pairCount / 4), (tunables.WORKGROUP_SIZE)), 65535)",
104
+ "y": "ceilDiv(ceilDiv((pairCount / 4), (tunables.WORKGROUP_SIZE)), 65535)",
105
  "z": 1
106
  }
107
  }
108
+ ],
109
+ "priority": 10
110
  },
111
  {
112
+ "id": "pairs",
 
113
  "derive": {
114
  "interleaved": "attrs.interleaved != 0",
115
  "mropeSectioned": "attrs.mrope_layout == 0",
 
116
  "scalar": "dtypes.T"
117
  },
118
  "passes": [
119
  {
120
  "id": "main",
121
+ "name": "MRotaryEmbedding",
122
  "shader": "mrotary-embedding.wgsl.jinja",
123
+ "derive": { "rank": "ranks.x" },
124
+ "bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params"],
125
  "dispatch": {
126
  "x": "min(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)",
127
  "y": "ceilDiv(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)",
build/webgpu/metadata.json CHANGED
@@ -1,21 +1,22 @@
1
  {
2
  "name": "com.microsoft.MRotaryEmbedding",
3
- "id": "_com_microsoft_mrotaryembedding_webgpu_83feafb",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "46hHXwN4zErrNbI5Kkvb00wSEt7LAjuixrnpiOUnuC8=",
11
- "manifest.json": "n2gNx3yadn8wpDXYWYTOGWFQVAiDogbtn03dXyQEAfE=",
12
- "mrotary-embedding.wgsl.jinja": "meXniZyv3KTe/4K0hd7HSYDjxS0ADty/sNevs/l3es8=",
13
- "test.json": "mQy2yFi3T6iEpqSliA+KGAipEQeV8moe0jnzSwH9ICI="
 
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
- "variants": { "rank3": ["mrotary-embedding.wgsl.jinja"], "rank4": ["mrotary-embedding.wgsl.jinja"] }
20
  }
21
  }
 
1
  {
2
  "name": "com.microsoft.MRotaryEmbedding",
3
+ "id": "_com_microsoft_mrotaryembedding_webgpu_9f14b3e",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "O2zNl+WongpRuCa+0uG2J/53uKtH0dOeHjs0V/SP5NQ=",
11
+ "manifest.json": "m6PieJLI+fSwIhyAV5llQg4PcQyE9OtvRC0w/Z3jOy0=",
12
+ "mrotary-embedding-vector.wgsl.jinja": "MZwbSQREHhZKFMGn1YC7CJATkbgtimGoPwdnzafMOn8=",
13
+ "mrotary-embedding.wgsl.jinja": "bMGNQWBiANeYtK1oQO0dRJjn/JwkB8ZTu/Zw/uKW2M4=",
14
+ "test.json": "TLcimbROUuwC6I8+3L/QcPyKX9brKS3g2eH68BRDfqI="
15
  }
16
  },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
  "webgpu": {
19
+ "manifestSpec": "2.1",
20
+ "variants": { "pairs4": ["mrotary-embedding-vector.wgsl.jinja"], "pairs": ["mrotary-embedding.wgsl.jinja"] }
21
  }
22
  }
build/webgpu/mrotary-embedding-vector.wgsl.jinja ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro mrope_cache_index(pair, suffix="") %}
2
+ // Which of the three position streams owns this cache column.
3
+ var stream{{ suffix }} = 0u;
4
+ {% if mropeSectioned %}
5
+ // Sectioned (Qwen2-VL, Qwen2.5-VL): three contiguous chunks [T | H | W].
6
+ if ({{ pair }} >= params.section0 && {{ pair }} < params.section0 + params.section1) {
7
+ stream{{ suffix }} = 1u;
8
+ } else if ({{ pair }} >= params.section0 + params.section1) {
9
+ stream{{ suffix }} = 2u;
10
+ }
11
+ {% else %}
12
+ // Interleaved (Qwen3-VL, Qwen3.5): T everywhere, then H over every third column
13
+ // from offset 1 and W over every third column from offset 2, each bounded by its
14
+ // own section length.
15
+ if ({{ pair }} % 3u == 1u && {{ pair }} / 3u < params.section1) {
16
+ stream{{ suffix }} = 1u;
17
+ } else if ({{ pair }} % 3u == 2u && {{ pair }} / 3u < params.section2) {
18
+ stream{{ suffix }} = 2u;
19
+ }
20
+ {% endif %}
21
+
22
+ let pos{{ suffix }} = position_ids[(stream{{ suffix }} * params.batchSize + batch) * params.sequenceLength + token];
23
+ let valid{{ suffix }} = pos{{ suffix }} >= 0 && u32(pos{{ suffix }}) < params.maxSequenceLength;
24
+ let cache{{ suffix }} = u32(pos{{ suffix }}) * params.halfRotaryDim + {{ pair }};
25
+ {% endmacro %}
26
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
27
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
28
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
29
+ // per-axis workgroup fold width.
30
+ {% if bound == "" %}
31
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% elif guardInline %}
32
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
33
+ if ({{ name }} >= {{ bound }}) { return; }{% else %}
34
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
35
+ if ({{ name }} >= {{ bound }}) {
36
+ return;
37
+ }{% endif %}{% endmacro %}
38
+ {{ env.wgsl.resourceDeclarations }}
39
+ // A contiguous vector owns several split-half rotation pairs. Each component
40
+ // gathers its own stream position, including vectors crossing section boundaries.
41
+ @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
42
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
43
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "p", "params.pairCount / " ~ pairWidth ~ "u") }}
44
+ let tasksPerHead = params.headSize / {{ 2 * pairWidth }}u;
45
+ let task = p % tasksPerHead;
46
+ let headFlat = p / tasksPerHead;
47
+ let base = headFlat * (params.headSize / {{ pairWidth }}u);
48
+ {% if rank == 3 %}
49
+ let token = (headFlat / params.numHeads) % params.sequenceLength;
50
+ let batch = headFlat / (params.numHeads * params.sequenceLength);
51
+ {% else %}
52
+ let token = headFlat % params.sequenceLength;
53
+ let batch = (headFlat / params.sequenceLength) / params.numHeads;
54
+ {% endif %}
55
+ {% for i in range(pairWidth) %}
56
+ let pair{{ i }} = task * {{ pairWidth }}u + {{ i }}u;
57
+ {{ mrope_cache_index("pair" ~ i, i) }}
58
+ {% endfor %}
59
+ let a = vec{{ pairWidth }}<f32>(x[base + task]);
60
+ let bIndex = base + params.halfRotaryDim / {{ pairWidth }}u + task;
61
+ let b = vec{{ pairWidth }}<f32>(x[bIndex]);
62
+ var outA = x[base + task];
63
+ var outB = x[bIndex];
64
+ {% for i in range(pairWidth) %}
65
+ if (valid{{ i }}) {
66
+ // Round scaled coefficients to storage T before evaluating the rotation.
67
+ let cf = f32({{ scalar }}(f32(cos_cache[cache{{ i }}]) * params.scale));
68
+ let sf = f32({{ scalar }}(f32(sin_cache[cache{{ i }}]) * params.scale));
69
+ outA[{{ i }}] = {{ scalar }}(a[{{ i }}] * cf - b[{{ i }}] * sf);
70
+ outB[{{ i }}] = {{ scalar }}(a[{{ i }}] * sf + b[{{ i }}] * cf);
71
+ }
72
+ {% endfor %}
73
+ y[base + task] = outA;
74
+ y[bIndex] = outB;
75
+ }
build/webgpu/mrotary-embedding.wgsl.jinja CHANGED
@@ -1,38 +1,37 @@
1
- {% macro flat_index_2d(name="i", bound="params.count", guardInline=false, note="dispatch-limit") %}
2
- {% if note == "dispatch-limit" %}
3
- // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
- // per-axis workgroup fold width (outputs > 16.7M elements).
5
- {% elif note == "limit" %}
6
- // 2D-folded flat index: gid.y carries the high bits past the dispatch's
7
- // per-axis workgroup fold width.
8
- {% elif note == "device-axis" %}
9
- // The flat dispatch is folded across x/y at a fixed per-axis workgroup
10
- // width; gid.y carries the high portion of the output index.
11
- {% elif note == "vec4-limit" %}
12
- // 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
13
- // per-axis workgroup fold width (the dispatch caps x and spills into y).
14
- {% elif note == "element-limit" %}
15
- // 2D-folded flat element index: gid.y carries the high bits past the
16
- // dispatch's per-axis workgroup fold width.
17
- {% elif note == "dispatch" %}
 
 
 
 
 
 
 
 
 
 
18
  // 2D-folded flat index: gid.y carries the high bits past the dispatch's
19
  // per-axis workgroup fold width.
20
- {% endif %}
21
- {% if bound == "" %}
22
- let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
23
- {%- elif guardInline %}
24
- let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
25
- if ({{ name }} >= {{ bound }}) { return; }
26
- {%- else %}
27
- let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
28
  if ({{ name }} >= {{ bound }}) {
29
  return;
30
- }
31
- {%- endif %}
32
- {% endmacro %}
33
-
34
- {% if usesF16 %}enable f16;
35
- {% endif %}{{ env.wgsl.resourceDeclarations }}
36
  // M-RoPE: three position streams (temporal, height, width) index one cos/sin cache.
37
  // The half-rotary axis is partitioned among the streams, so each rotated pair picks
38
  // its own stream and gathers that stream's position for this token. One invocation
@@ -40,7 +39,7 @@
40
  // use the remaining pair lanes to copy two unchanged tail values.
41
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
42
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
43
- {{ flat_index_2d("p", "params.pairCount", note="") }}
44
 
45
  let tasksPerHead = (params.headSize + 1u) / 2u;
46
  let task = p % tasksPerHead;
@@ -68,32 +67,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
68
  let batch = (headFlat / params.sequenceLength) / params.numHeads;
69
  {% endif %}
70
 
71
- // Which of the three position streams owns this cache column.
72
- var stream = 0u;
73
- {% if mropeSectioned %}
74
- // Sectioned (Qwen2-VL, Qwen2.5-VL): three contiguous chunks [T | H | W].
75
- if (pair >= params.section0 && pair < params.section0 + params.section1) {
76
- stream = 1u;
77
- } else if (pair >= params.section0 + params.section1) {
78
- stream = 2u;
79
- }
80
- {% else %}
81
- // Interleaved (Qwen3-VL, Qwen3.5): T everywhere, then H over every third column
82
- // from offset 1 and W over every third column from offset 2, each bounded by its
83
- // own section length.
84
- if (pair % 3u == 1u && pair / 3u < params.section1) {
85
- stream = 1u;
86
- } else if (pair % 3u == 2u && pair / 3u < params.section2) {
87
- stream = 2u;
88
- }
89
- {% endif %}
90
-
91
- let pos = position_ids[(stream * params.batchSize + batch) * params.sequenceLength + token];
92
- let cache = pos * params.halfRotaryDim + pair;
93
- // Round scaled cache values to T before rotation. This intermediate cast is
94
- // observable for f16 and must not be deferred to the output store.
95
- let cf = f32({{ scalar }}(f32(cos_cache[cache]) * params.scale));
96
- let sf = f32({{ scalar }}(f32(sin_cache[cache]) * params.scale));
97
  {% if interleaved %}
98
  let aOffset = pair * 2u;
99
  let bOffset = aOffset + 1u;
@@ -101,6 +75,15 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
101
  let aOffset = pair;
102
  let bOffset = pair + params.halfRotaryDim;
103
  {% endif %}
 
 
 
 
 
 
 
 
 
104
  let a = f32(x[base + aOffset]);
105
  let b = f32(x[base + bOffset]);
106
  y[base + aOffset] = {{ scalar }}(a * cf - b * sf);
 
1
+ {% macro mrope_cache_index(pair, suffix="") %}
2
+ // Which of the three position streams owns this cache column.
3
+ var stream{{ suffix }} = 0u;
4
+ {% if mropeSectioned %}
5
+ // Sectioned (Qwen2-VL, Qwen2.5-VL): three contiguous chunks [T | H | W].
6
+ if ({{ pair }} >= params.section0 && {{ pair }} < params.section0 + params.section1) {
7
+ stream{{ suffix }} = 1u;
8
+ } else if ({{ pair }} >= params.section0 + params.section1) {
9
+ stream{{ suffix }} = 2u;
10
+ }
11
+ {% else %}
12
+ // Interleaved (Qwen3-VL, Qwen3.5): T everywhere, then H over every third column
13
+ // from offset 1 and W over every third column from offset 2, each bounded by its
14
+ // own section length.
15
+ if ({{ pair }} % 3u == 1u && {{ pair }} / 3u < params.section1) {
16
+ stream{{ suffix }} = 1u;
17
+ } else if ({{ pair }} % 3u == 2u && {{ pair }} / 3u < params.section2) {
18
+ stream{{ suffix }} = 2u;
19
+ }
20
+ {% endif %}
21
+
22
+ let pos{{ suffix }} = position_ids[(stream{{ suffix }} * params.batchSize + batch) * params.sequenceLength + token];
23
+ let valid{{ suffix }} = pos{{ suffix }} >= 0 && u32(pos{{ suffix }}) < params.maxSequenceLength;
24
+ let cache{{ suffix }} = u32(pos{{ suffix }}) * params.halfRotaryDim + {{ pair }};
25
+ {% endmacro %}
26
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
27
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
28
  // 2D-folded flat index: gid.y carries the high bits past the dispatch's
29
  // per-axis workgroup fold width.
30
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
 
 
 
 
 
 
 
31
  if ({{ name }} >= {{ bound }}) {
32
  return;
33
+ }{% endmacro %}
34
+ {{ env.wgsl.resourceDeclarations }}
 
 
 
 
35
  // M-RoPE: three position streams (temporal, height, width) index one cos/sin cache.
36
  // The half-rotary axis is partitioned among the streams, so each rotated pair picks
37
  // its own stream and gathers that stream's position for this token. One invocation
 
39
  // use the remaining pair lanes to copy two unchanged tail values.
40
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
41
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
42
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "p", "params.pairCount") }}
43
 
44
  let tasksPerHead = (params.headSize + 1u) / 2u;
45
  let task = p % tasksPerHead;
 
67
  let batch = (headFlat / params.sequenceLength) / params.numHeads;
68
  {% endif %}
69
 
70
+ {{ mrope_cache_index("pair") }}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
71
  {% if interleaved %}
72
  let aOffset = pair * 2u;
73
  let bOffset = aOffset + 1u;
 
75
  let aOffset = pair;
76
  let bOffset = pair + params.halfRotaryDim;
77
  {% endif %}
78
+ if (!valid) {
79
+ y[base + aOffset] = x[base + aOffset];
80
+ y[base + bOffset] = x[base + bOffset];
81
+ return;
82
+ }
83
+ // Round scaled cache values to T before rotation. This intermediate cast is
84
+ // observable for f16 and must not be deferred to the output store.
85
+ let cf = f32({{ scalar }}(f32(cos_cache[cache]) * params.scale));
86
+ let sf = f32({{ scalar }}(f32(sin_cache[cache]) * params.scale));
87
  let a = f32(x[base + aOffset]);
88
  let b = f32(x[base + bOffset]);
89
  y[base + aOffset] = {{ scalar }}(a * cf - b * sf);
build/webgpu/test.json CHANGED
@@ -1,6 +1,8 @@
1
  {
2
  "fixtureArrays": {
3
- "ort_sectioned_rank3_input_x": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24]
 
 
4
  },
5
  "cases": [
6
  {
@@ -19,7 +21,7 @@
19
  "inputs": {
20
  "x": { "dtype": "float32", "shape": [1, 1, 2, 5], "data": { "kind": "linspace", "start": -1.25, "end": 1.5 } },
21
  "positionIds": {
22
- "dtype": "uint32",
23
  "shape": [3, 1, 2],
24
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 0] }
25
  },
@@ -50,7 +52,7 @@
50
  "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_input_x" } }
51
  },
52
  "positionIds": {
53
- "dtype": "uint32",
54
  "shape": [3, 1, 2],
55
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3] }
56
  },
@@ -70,10 +72,7 @@
70
  "dtype": "float32",
71
  "shape": [1, 2, 12],
72
  "tolerance": 0.00001,
73
- "data": {
74
- "kind": "values",
75
- "values": [0.6, 1.42, 2.34, 4.1, 5.87, 7.98, 6.0, 7.12, 8.34, 10.7, 13.49, 16.62, 11.9, 13.37, 14.94, 19.55, 23.51, 27.81, 17.6, 19.37, 21.24, 27.05, 32.03, 37.35]
76
- }
77
  }
78
  }
79
  },
@@ -99,7 +98,7 @@
99
  "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_input_x" } }
100
  },
101
  "positionIds": {
102
- "dtype": "uint32",
103
  "shape": [3, 1, 2],
104
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3] }
105
  },
@@ -119,10 +118,7 @@
119
  "dtype": "float32",
120
  "shape": [1, 2, 2, 6],
121
  "tolerance": 0.00001,
122
- "data": {
123
- "kind": "values",
124
- "values": [0.4, 1.05, 1.345, 2.46, 2.39, 4.21, 3.25, 4.925, 4.395, 6.995, 5.64, 9.405, 5.8, 7.65, 7.045, 10.08, 8.39, 12.85, 8.95, 12.425, 10.395, 15.515, 11.94, 18.945]
125
- }
126
  }
127
  }
128
  },
@@ -142,7 +138,7 @@
142
  "inputs": {
143
  "x": { "dtype": "float32", "shape": [2, 3, 32], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
144
  "positionIds": {
145
- "dtype": "uint32",
146
  "shape": [3, 2, 3],
147
  "data": { "kind": "values", "values": [0, 1, 2, 1, 2, 3, 2, 3, 4, 3, 4, 5, 4, 5, 6, 5, 6, 7] }
148
  },
@@ -167,7 +163,7 @@
167
  "inputs": {
168
  "x": { "dtype": "float32", "shape": [1, 2, 4, 12], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
169
  "positionIds": {
170
- "dtype": "uint32",
171
  "shape": [3, 1, 4],
172
  "data": { "kind": "values", "values": [0, 1, 2, 3, 1, 2, 3, 4, 2, 3, 4, 5] }
173
  },
@@ -185,7 +181,7 @@
185
  "inputs": {
186
  "x": { "dtype": "float32", "shape": [1, 4, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
187
  "positionIds": {
188
- "dtype": "uint32",
189
  "shape": [3, 1, 4],
190
  "data": { "kind": "values", "values": [0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2, 3] }
191
  },
@@ -209,7 +205,7 @@
209
  "inputs": {
210
  "x": { "dtype": "float16", "shape": [1, 2, 3, 12], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
211
  "positionIds": {
212
- "dtype": "uint32",
213
  "shape": [3, 1, 3],
214
  "data": { "kind": "values", "values": [0, 1, 2, 1, 2, 3, 2, 3, 4] }
215
  },
@@ -232,7 +228,7 @@
232
  "inputs": {
233
  "x": { "dtype": "float32", "shape": [2, 2, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
234
  "positionIds": {
235
- "dtype": "uint32",
236
  "shape": [3, 2, 2],
237
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3, 3, 4, 4, 0, 0, 1] }
238
  },
@@ -240,6 +236,1048 @@
240
  "sin": { "dtype": "float32", "shape": [5, 4], "data": { "kind": "linspace", "start": 1.0, "end": -1.0 } }
241
  },
242
  "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 16], "tolerance": 0.00001, "relTolerance": 0.00001 } }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
243
  }
244
  ]
245
  }
 
1
  {
2
  "fixtureArrays": {
3
+ "ort_sectioned_rank3_input_x": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24],
4
+ "ort_sectioned_rank3_output_y": [0.6, 1.42, 2.34, 4.1, 5.87, 7.98, 6, 7.12, 8.34, 10.7, 13.49, 16.62, 11.9, 13.37, 14.94, 19.55, 23.51, 27.81, 17.6, 19.37, 21.24, 27.05, 32.03, 37.35],
5
+ "ort_interleaved_rank4_output_y": [0.4, 1.05, 1.345, 2.46, 2.39, 4.21, 3.25, 4.925, 4.395, 6.995, 5.64, 9.405, 5.8, 7.65, 7.045, 10.08, 8.39, 12.85, 8.95, 12.425, 10.395, 15.515, 11.94, 18.945]
6
  },
7
  "cases": [
8
  {
 
21
  "inputs": {
22
  "x": { "dtype": "float32", "shape": [1, 1, 2, 5], "data": { "kind": "linspace", "start": -1.25, "end": 1.5 } },
23
  "positionIds": {
24
+ "dtype": "int32",
25
  "shape": [3, 1, 2],
26
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 0] }
27
  },
 
52
  "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_input_x" } }
53
  },
54
  "positionIds": {
55
+ "dtype": "int32",
56
  "shape": [3, 1, 2],
57
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3] }
58
  },
 
72
  "dtype": "float32",
73
  "shape": [1, 2, 12],
74
  "tolerance": 0.00001,
75
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_output_y" } }
 
 
 
76
  }
77
  }
78
  },
 
98
  "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_input_x" } }
99
  },
100
  "positionIds": {
101
+ "dtype": "int32",
102
  "shape": [3, 1, 2],
103
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3] }
104
  },
 
118
  "dtype": "float32",
119
  "shape": [1, 2, 2, 6],
120
  "tolerance": 0.00001,
121
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_interleaved_rank4_output_y" } }
 
 
 
122
  }
123
  }
124
  },
 
138
  "inputs": {
139
  "x": { "dtype": "float32", "shape": [2, 3, 32], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
140
  "positionIds": {
141
+ "dtype": "int32",
142
  "shape": [3, 2, 3],
143
  "data": { "kind": "values", "values": [0, 1, 2, 1, 2, 3, 2, 3, 4, 3, 4, 5, 4, 5, 6, 5, 6, 7] }
144
  },
 
163
  "inputs": {
164
  "x": { "dtype": "float32", "shape": [1, 2, 4, 12], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
165
  "positionIds": {
166
+ "dtype": "int32",
167
  "shape": [3, 1, 4],
168
  "data": { "kind": "values", "values": [0, 1, 2, 3, 1, 2, 3, 4, 2, 3, 4, 5] }
169
  },
 
181
  "inputs": {
182
  "x": { "dtype": "float32", "shape": [1, 4, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
183
  "positionIds": {
184
+ "dtype": "int32",
185
  "shape": [3, 1, 4],
186
  "data": { "kind": "values", "values": [0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2, 3] }
187
  },
 
205
  "inputs": {
206
  "x": { "dtype": "float16", "shape": [1, 2, 3, 12], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
207
  "positionIds": {
208
+ "dtype": "int32",
209
  "shape": [3, 1, 3],
210
  "data": { "kind": "values", "values": [0, 1, 2, 1, 2, 3, 2, 3, 4] }
211
  },
 
228
  "inputs": {
229
  "x": { "dtype": "float32", "shape": [2, 2, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
230
  "positionIds": {
231
+ "dtype": "int32",
232
  "shape": [3, 2, 2],
233
  "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3, 3, 4, 4, 0, 0, 1] }
234
  },
 
236
  "sin": { "dtype": "float32", "shape": [5, 4], "data": { "kind": "linspace", "start": 1.0, "end": -1.0 } }
237
  },
238
  "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 16], "tolerance": 0.00001, "relTolerance": 0.00001 } }
239
+ },
240
+ {
241
+ "name": "vector_boundary_float32_rank3_d8_layout0_zero1",
242
+ "attrs": {
243
+ "num_heads": 3,
244
+ "rotary_embedding_dim": 8,
245
+ "interleaved": 0,
246
+ "mrope_layout": 0,
247
+ "mrope_section": [0, 3, 1],
248
+ "scale": 0.3
249
+ },
250
+ "inputs": {
251
+ "x": {
252
+ "dtype": "float32",
253
+ "shape": [2, 5, 24],
254
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
255
+ },
256
+ "positionIds": {
257
+ "dtype": "int32",
258
+ "shape": [3, 2, 5],
259
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
260
+ },
261
+ "cos": {
262
+ "dtype": "float32",
263
+ "shape": [5, 4],
264
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
265
+ },
266
+ "sin": {
267
+ "dtype": "float32",
268
+ "shape": [5, 4],
269
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
270
+ }
271
+ },
272
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 5, 24], "tolerance": 0.00001, "relTolerance": 0.00001 } },
273
+ "provenance": {
274
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
275
+ }
276
+ },
277
+ {
278
+ "name": "vector_boundary_float32_rank3_d16_layout0_zero0",
279
+ "attrs": {
280
+ "num_heads": 3,
281
+ "rotary_embedding_dim": 16,
282
+ "interleaved": 0,
283
+ "mrope_layout": 0,
284
+ "mrope_section": [1, 2, 5],
285
+ "scale": 0.3
286
+ },
287
+ "inputs": {
288
+ "x": {
289
+ "dtype": "float32",
290
+ "shape": [2, 5, 48],
291
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
292
+ },
293
+ "positionIds": {
294
+ "dtype": "int32",
295
+ "shape": [3, 2, 5],
296
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
297
+ },
298
+ "cos": {
299
+ "dtype": "float32",
300
+ "shape": [5, 8],
301
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
302
+ },
303
+ "sin": {
304
+ "dtype": "float32",
305
+ "shape": [5, 8],
306
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
307
+ }
308
+ },
309
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 5, 48], "tolerance": 0.00001, "relTolerance": 0.00001 } },
310
+ "provenance": {
311
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
312
+ }
313
+ },
314
+ {
315
+ "name": "vector_boundary_float32_rank3_d24_layout0_zero0",
316
+ "attrs": {
317
+ "num_heads": 3,
318
+ "rotary_embedding_dim": 24,
319
+ "interleaved": 0,
320
+ "mrope_layout": 0,
321
+ "mrope_section": [1, 2, 9],
322
+ "scale": 0.3
323
+ },
324
+ "inputs": {
325
+ "x": {
326
+ "dtype": "float32",
327
+ "shape": [2, 5, 72],
328
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
329
+ },
330
+ "positionIds": {
331
+ "dtype": "int32",
332
+ "shape": [3, 2, 5],
333
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
334
+ },
335
+ "cos": {
336
+ "dtype": "float32",
337
+ "shape": [5, 12],
338
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
339
+ },
340
+ "sin": {
341
+ "dtype": "float32",
342
+ "shape": [5, 12],
343
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
344
+ }
345
+ },
346
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 5, 72], "tolerance": 0.00001, "relTolerance": 0.00001 } },
347
+ "provenance": {
348
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
349
+ }
350
+ },
351
+ {
352
+ "name": "vector_boundary_float32_rank4_d8_layout1_zero1",
353
+ "attrs": {
354
+ "num_heads": 3,
355
+ "rotary_embedding_dim": 8,
356
+ "interleaved": 0,
357
+ "mrope_layout": 1,
358
+ "mrope_section": [0, 3, 1],
359
+ "scale": 0.3
360
+ },
361
+ "inputs": {
362
+ "x": {
363
+ "dtype": "float32",
364
+ "shape": [2, 3, 5, 8],
365
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
366
+ },
367
+ "positionIds": {
368
+ "dtype": "int32",
369
+ "shape": [3, 2, 5],
370
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
371
+ },
372
+ "cos": {
373
+ "dtype": "float32",
374
+ "shape": [5, 4],
375
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
376
+ },
377
+ "sin": {
378
+ "dtype": "float32",
379
+ "shape": [5, 4],
380
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
381
+ }
382
+ },
383
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 5, 8], "tolerance": 0.00001, "relTolerance": 0.00001 } },
384
+ "provenance": {
385
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
386
+ }
387
+ },
388
+ {
389
+ "name": "vector_boundary_float32_rank4_d16_layout1_zero0",
390
+ "attrs": {
391
+ "num_heads": 3,
392
+ "rotary_embedding_dim": 16,
393
+ "interleaved": 0,
394
+ "mrope_layout": 1,
395
+ "mrope_section": [1, 2, 5],
396
+ "scale": 0.3
397
+ },
398
+ "inputs": {
399
+ "x": {
400
+ "dtype": "float32",
401
+ "shape": [2, 3, 5, 16],
402
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
403
+ },
404
+ "positionIds": {
405
+ "dtype": "int32",
406
+ "shape": [3, 2, 5],
407
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
408
+ },
409
+ "cos": {
410
+ "dtype": "float32",
411
+ "shape": [5, 8],
412
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
413
+ },
414
+ "sin": {
415
+ "dtype": "float32",
416
+ "shape": [5, 8],
417
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
418
+ }
419
+ },
420
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 5, 16], "tolerance": 0.00001, "relTolerance": 0.00001 } },
421
+ "provenance": {
422
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
423
+ }
424
+ },
425
+ {
426
+ "name": "vector_boundary_float32_rank4_d24_layout1_zero0",
427
+ "attrs": {
428
+ "num_heads": 3,
429
+ "rotary_embedding_dim": 24,
430
+ "interleaved": 0,
431
+ "mrope_layout": 1,
432
+ "mrope_section": [1, 2, 9],
433
+ "scale": 0.3
434
+ },
435
+ "inputs": {
436
+ "x": {
437
+ "dtype": "float32",
438
+ "shape": [2, 3, 5, 24],
439
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
440
+ },
441
+ "positionIds": {
442
+ "dtype": "int32",
443
+ "shape": [3, 2, 5],
444
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
445
+ },
446
+ "cos": {
447
+ "dtype": "float32",
448
+ "shape": [5, 12],
449
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
450
+ },
451
+ "sin": {
452
+ "dtype": "float32",
453
+ "shape": [5, 12],
454
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
455
+ }
456
+ },
457
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 5, 24], "tolerance": 0.00001, "relTolerance": 0.00001 } },
458
+ "provenance": {
459
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
460
+ }
461
+ },
462
+ {
463
+ "name": "vector_boundary_float16_rank3_d8_layout0_zero1",
464
+ "attrs": {
465
+ "num_heads": 3,
466
+ "rotary_embedding_dim": 8,
467
+ "interleaved": 0,
468
+ "mrope_layout": 0,
469
+ "mrope_section": [0, 3, 1],
470
+ "scale": 0.3
471
+ },
472
+ "inputs": {
473
+ "x": {
474
+ "dtype": "float16",
475
+ "shape": [2, 5, 24],
476
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
477
+ },
478
+ "positionIds": {
479
+ "dtype": "int32",
480
+ "shape": [3, 2, 5],
481
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
482
+ },
483
+ "cos": {
484
+ "dtype": "float16",
485
+ "shape": [5, 4],
486
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
487
+ },
488
+ "sin": {
489
+ "dtype": "float16",
490
+ "shape": [5, 4],
491
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
492
+ }
493
+ },
494
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 5, 24], "tolerance": 0.01, "relTolerance": 0.01 } },
495
+ "provenance": {
496
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
497
+ }
498
+ },
499
+ {
500
+ "name": "vector_boundary_float16_rank3_d16_layout0_zero0",
501
+ "attrs": {
502
+ "num_heads": 3,
503
+ "rotary_embedding_dim": 16,
504
+ "interleaved": 0,
505
+ "mrope_layout": 0,
506
+ "mrope_section": [1, 2, 5],
507
+ "scale": 0.3
508
+ },
509
+ "inputs": {
510
+ "x": {
511
+ "dtype": "float16",
512
+ "shape": [2, 5, 48],
513
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
514
+ },
515
+ "positionIds": {
516
+ "dtype": "int32",
517
+ "shape": [3, 2, 5],
518
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
519
+ },
520
+ "cos": {
521
+ "dtype": "float16",
522
+ "shape": [5, 8],
523
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
524
+ },
525
+ "sin": {
526
+ "dtype": "float16",
527
+ "shape": [5, 8],
528
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
529
+ }
530
+ },
531
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 5, 48], "tolerance": 0.01, "relTolerance": 0.01 } },
532
+ "provenance": {
533
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
534
+ }
535
+ },
536
+ {
537
+ "name": "vector_boundary_float16_rank3_d24_layout0_zero0",
538
+ "attrs": {
539
+ "num_heads": 3,
540
+ "rotary_embedding_dim": 24,
541
+ "interleaved": 0,
542
+ "mrope_layout": 0,
543
+ "mrope_section": [1, 2, 9],
544
+ "scale": 0.3
545
+ },
546
+ "inputs": {
547
+ "x": {
548
+ "dtype": "float16",
549
+ "shape": [2, 5, 72],
550
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
551
+ },
552
+ "positionIds": {
553
+ "dtype": "int32",
554
+ "shape": [3, 2, 5],
555
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
556
+ },
557
+ "cos": {
558
+ "dtype": "float16",
559
+ "shape": [5, 12],
560
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
561
+ },
562
+ "sin": {
563
+ "dtype": "float16",
564
+ "shape": [5, 12],
565
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
566
+ }
567
+ },
568
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 5, 72], "tolerance": 0.01, "relTolerance": 0.01 } },
569
+ "provenance": {
570
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
571
+ }
572
+ },
573
+ {
574
+ "name": "vector_boundary_float16_rank4_d8_layout1_zero1",
575
+ "attrs": {
576
+ "num_heads": 3,
577
+ "rotary_embedding_dim": 8,
578
+ "interleaved": 0,
579
+ "mrope_layout": 1,
580
+ "mrope_section": [0, 3, 1],
581
+ "scale": 0.3
582
+ },
583
+ "inputs": {
584
+ "x": {
585
+ "dtype": "float16",
586
+ "shape": [2, 3, 5, 8],
587
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
588
+ },
589
+ "positionIds": {
590
+ "dtype": "int32",
591
+ "shape": [3, 2, 5],
592
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
593
+ },
594
+ "cos": {
595
+ "dtype": "float16",
596
+ "shape": [5, 4],
597
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
598
+ },
599
+ "sin": {
600
+ "dtype": "float16",
601
+ "shape": [5, 4],
602
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
603
+ }
604
+ },
605
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 3, 5, 8], "tolerance": 0.01, "relTolerance": 0.01 } },
606
+ "provenance": {
607
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
608
+ }
609
+ },
610
+ {
611
+ "name": "vector_boundary_float16_rank4_d16_layout1_zero0",
612
+ "attrs": {
613
+ "num_heads": 3,
614
+ "rotary_embedding_dim": 16,
615
+ "interleaved": 0,
616
+ "mrope_layout": 1,
617
+ "mrope_section": [1, 2, 5],
618
+ "scale": 0.3
619
+ },
620
+ "inputs": {
621
+ "x": {
622
+ "dtype": "float16",
623
+ "shape": [2, 3, 5, 16],
624
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
625
+ },
626
+ "positionIds": {
627
+ "dtype": "int32",
628
+ "shape": [3, 2, 5],
629
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
630
+ },
631
+ "cos": {
632
+ "dtype": "float16",
633
+ "shape": [5, 8],
634
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
635
+ },
636
+ "sin": {
637
+ "dtype": "float16",
638
+ "shape": [5, 8],
639
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
640
+ }
641
+ },
642
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 3, 5, 16], "tolerance": 0.01, "relTolerance": 0.01 } },
643
+ "provenance": {
644
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
645
+ }
646
+ },
647
+ {
648
+ "name": "vector_boundary_float16_rank4_d24_layout1_zero0",
649
+ "attrs": {
650
+ "num_heads": 3,
651
+ "rotary_embedding_dim": 24,
652
+ "interleaved": 0,
653
+ "mrope_layout": 1,
654
+ "mrope_section": [1, 2, 9],
655
+ "scale": 0.3
656
+ },
657
+ "inputs": {
658
+ "x": {
659
+ "dtype": "float16",
660
+ "shape": [2, 3, 5, 24],
661
+ "data": { "kind": "fillFloat32", "sinStep": 0.113, "cosStep": 0.317, "scale": 3.0 }
662
+ },
663
+ "positionIds": {
664
+ "dtype": "int32",
665
+ "shape": [3, 2, 5],
666
+ "data": { "kind": "cycle", "values": [0, 4, 2, 1, 3, 4, 0] }
667
+ },
668
+ "cos": {
669
+ "dtype": "float16",
670
+ "shape": [5, 12],
671
+ "data": { "kind": "fillFloat32", "sinStep": 0.213, "cosStep": 0.437, "scale": 0.71 }
672
+ },
673
+ "sin": {
674
+ "dtype": "float16",
675
+ "shape": [5, 12],
676
+ "data": { "kind": "fillFloat32", "sinStep": 0.619, "cosStep": 0.137, "scale": 0.83 }
677
+ }
678
+ },
679
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 3, 5, 24], "tolerance": 0.01, "relTolerance": 0.01 } },
680
+ "provenance": {
681
+ "notes": "Aligned vector pairs with a partial final workgroup, section-crossing cache columns, multiple batches, non-unit scale and both tensor ranks. Zero temporal section is covered at the smallest head width."
682
+ }
683
+ },
684
+ {
685
+ "name": "ort_partial_rotary_dim_copies_tail_rank3",
686
+ "attrs": {
687
+ "num_heads": 1,
688
+ "rotary_embedding_dim": 6,
689
+ "mrope_layout": 0,
690
+ "interleaved": 0,
691
+ "scale": 1,
692
+ "mrope_section": [1, 1, 1]
693
+ },
694
+ "inputs": {
695
+ "x": {
696
+ "dtype": "float32",
697
+ "shape": [1, 1, 8],
698
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0] }
699
+ },
700
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [0, 0, 0] } },
701
+ "cos": {
702
+ "dtype": "float32",
703
+ "shape": [4, 3],
704
+ "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0] }
705
+ },
706
+ "sin": {
707
+ "dtype": "float32",
708
+ "shape": [4, 3],
709
+ "data": { "kind": "values", "values": [1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0] }
710
+ }
711
+ },
712
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 1, 8], "tolerance": 0.000001, "relTolerance": 0.000001 } }
713
+ },
714
+ {
715
+ "name": "ort_empty_rank3_input",
716
+ "attrs": { "num_heads": 1, "mrope_layout": 0, "interleaved": 0, "mrope_section": [0, 0, 0] },
717
+ "inputs": {
718
+ "x": { "dtype": "float32", "shape": [1, 1, 0], "data": { "kind": "values", "values": [] } },
719
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [0, 0, 0] } },
720
+ "cos": { "dtype": "float32", "shape": [1, 0], "data": { "kind": "values", "values": [] } },
721
+ "sin": { "dtype": "float32", "shape": [1, 0], "data": { "kind": "values", "values": [] } }
722
+ },
723
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 1, 0], "tolerance": 0.000001, "relTolerance": 0.000001 } }
724
+ },
725
+ {
726
+ "name": "ort_empty_rank4_input",
727
+ "attrs": {
728
+ "num_heads": 2,
729
+ "rotary_embedding_dim": 6,
730
+ "mrope_layout": 0,
731
+ "interleaved": 0,
732
+ "mrope_section": [1, 1, 1]
733
+ },
734
+ "inputs": {
735
+ "x": { "dtype": "float32", "shape": [1, 2, 0, 6], "data": { "kind": "values", "values": [] } },
736
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 0], "data": { "kind": "values", "values": [] } },
737
+ "cos": { "dtype": "float32", "shape": [1, 3], "data": { "kind": "values", "values": [1.0, 1.0, 1.0] } },
738
+ "sin": { "dtype": "float32", "shape": [1, 3], "data": { "kind": "values", "values": [0.0, 0.0, 0.0] } }
739
+ },
740
+ "outputs": { "y": { "dtype": "float32", "shape": [1, 2, 0, 6], "tolerance": 0.000001, "relTolerance": 0.000001 } }
741
+ },
742
+ {
743
+ "name": "ort_sectioned_rank3_f16",
744
+ "attrs": {
745
+ "num_heads": 2,
746
+ "rotary_embedding_dim": 6,
747
+ "mrope_layout": 0,
748
+ "interleaved": 0,
749
+ "scale": 1,
750
+ "mrope_section": [1, 1, 1]
751
+ },
752
+ "inputs": {
753
+ "x": {
754
+ "dtype": "float16",
755
+ "shape": [1, 2, 12],
756
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_input_x" } }
757
+ },
758
+ "positionIds": {
759
+ "dtype": "int32",
760
+ "shape": [3, 1, 2],
761
+ "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3] }
762
+ },
763
+ "cos": {
764
+ "dtype": "float16",
765
+ "shape": [4, 3],
766
+ "data": { "kind": "values", "values": [1.0, 1.01, 1.02, 1.1, 1.11, 1.12, 1.2, 1.21, 1.22, 1.3, 1.31, 1.32] }
767
+ },
768
+ "sin": {
769
+ "dtype": "float16",
770
+ "shape": [4, 3],
771
+ "data": { "kind": "values", "values": [0.1, 0.11, 0.12, 0.15, 0.16, 0.17, 0.2, 0.21, 0.22, 0.25, 0.26, 0.27] }
772
+ }
773
+ },
774
+ "outputs": {
775
+ "y": {
776
+ "dtype": "float16",
777
+ "shape": [1, 2, 12],
778
+ "tolerance": 0.01,
779
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_output_y" } },
780
+ "relTolerance": 0.01
781
+ }
782
+ }
783
+ },
784
+ {
785
+ "name": "ort_interleaved_rank4_f16",
786
+ "attrs": {
787
+ "num_heads": 2,
788
+ "rotary_embedding_dim": 6,
789
+ "mrope_layout": 1,
790
+ "interleaved": 1,
791
+ "scale": 0.5,
792
+ "mrope_section": [1, 1, 1]
793
+ },
794
+ "inputs": {
795
+ "x": {
796
+ "dtype": "float16",
797
+ "shape": [1, 2, 2, 6],
798
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_sectioned_rank3_input_x" } }
799
+ },
800
+ "positionIds": {
801
+ "dtype": "int32",
802
+ "shape": [3, 1, 2],
803
+ "data": { "kind": "values", "values": [0, 1, 1, 2, 2, 3] }
804
+ },
805
+ "cos": {
806
+ "dtype": "float16",
807
+ "shape": [4, 3],
808
+ "data": { "kind": "values", "values": [1.0, 1.01, 1.02, 1.1, 1.11, 1.12, 1.2, 1.21, 1.22, 1.3, 1.31, 1.32] }
809
+ },
810
+ "sin": {
811
+ "dtype": "float16",
812
+ "shape": [4, 3],
813
+ "data": { "kind": "values", "values": [0.1, 0.11, 0.12, 0.15, 0.16, 0.17, 0.2, 0.21, 0.22, 0.25, 0.26, 0.27] }
814
+ }
815
+ },
816
+ "outputs": {
817
+ "y": {
818
+ "dtype": "float16",
819
+ "shape": [1, 2, 2, 6],
820
+ "tolerance": 0.01,
821
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_interleaved_rank4_output_y" } },
822
+ "relTolerance": 0.01
823
+ }
824
+ }
825
+ },
826
+ {
827
+ "name": "ort_oob_negative_stream0",
828
+ "provenance": {
829
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
830
+ "test": "PositionIdsOOBPassthroughAllStreams",
831
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
832
+ },
833
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
834
+ "inputs": {
835
+ "x": {
836
+ "dtype": "float32",
837
+ "shape": [1, 1, 1, 6],
838
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
839
+ },
840
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [-1, 0, 0] } },
841
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
842
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
843
+ },
844
+ "outputs": {
845
+ "y": {
846
+ "dtype": "float32",
847
+ "shape": [1, 1, 1, 6],
848
+ "data": { "kind": "values", "values": [1.0, -5.0, -6.0, 4.0, 2.0, 3.0] }
849
+ }
850
+ }
851
+ },
852
+ {
853
+ "name": "ort_oob_upper_bound_stream0",
854
+ "provenance": {
855
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
856
+ "test": "PositionIdsOOBPassthroughAllStreams",
857
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
858
+ },
859
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
860
+ "inputs": {
861
+ "x": {
862
+ "dtype": "float32",
863
+ "shape": [1, 1, 1, 6],
864
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
865
+ },
866
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [2, 0, 0] } },
867
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
868
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
869
+ },
870
+ "outputs": {
871
+ "y": {
872
+ "dtype": "float32",
873
+ "shape": [1, 1, 1, 6],
874
+ "data": { "kind": "values", "values": [1.0, -5.0, -6.0, 4.0, 2.0, 3.0] }
875
+ }
876
+ }
877
+ },
878
+ {
879
+ "name": "ort_oob_saturated_high_stream0",
880
+ "provenance": {
881
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
882
+ "test": "PositionIdsOOBPassthroughAllStreams",
883
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
884
+ },
885
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
886
+ "inputs": {
887
+ "x": {
888
+ "dtype": "float32",
889
+ "shape": [1, 1, 1, 6],
890
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
891
+ },
892
+ "positionIds": {
893
+ "dtype": "int32",
894
+ "shape": [3, 1, 1],
895
+ "data": { "kind": "values", "values": [2147483647, 0, 0] }
896
+ },
897
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
898
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
899
+ },
900
+ "outputs": {
901
+ "y": {
902
+ "dtype": "float32",
903
+ "shape": [1, 1, 1, 6],
904
+ "data": { "kind": "values", "values": [1.0, -5.0, -6.0, 4.0, 2.0, 3.0] }
905
+ }
906
+ }
907
+ },
908
+ {
909
+ "name": "ort_oob_saturated_low_stream0",
910
+ "provenance": {
911
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
912
+ "test": "PositionIdsOOBPassthroughAllStreams",
913
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
914
+ },
915
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
916
+ "inputs": {
917
+ "x": {
918
+ "dtype": "float32",
919
+ "shape": [1, 1, 1, 6],
920
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
921
+ },
922
+ "positionIds": {
923
+ "dtype": "int32",
924
+ "shape": [3, 1, 1],
925
+ "data": { "kind": "values", "values": [-2147483648, 0, 0] }
926
+ },
927
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
928
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
929
+ },
930
+ "outputs": {
931
+ "y": {
932
+ "dtype": "float32",
933
+ "shape": [1, 1, 1, 6],
934
+ "data": { "kind": "values", "values": [1.0, -5.0, -6.0, 4.0, 2.0, 3.0] }
935
+ }
936
+ }
937
+ },
938
+ {
939
+ "name": "ort_oob_negative_stream1",
940
+ "provenance": {
941
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
942
+ "test": "PositionIdsOOBPassthroughAllStreams",
943
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
944
+ },
945
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
946
+ "inputs": {
947
+ "x": {
948
+ "dtype": "float32",
949
+ "shape": [1, 1, 1, 6],
950
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
951
+ },
952
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [0, -1, 0] } },
953
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
954
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
955
+ },
956
+ "outputs": {
957
+ "y": {
958
+ "dtype": "float32",
959
+ "shape": [1, 1, 1, 6],
960
+ "data": { "kind": "values", "values": [-4.0, 2.0, -6.0, 1.0, 5.0, 3.0] }
961
+ }
962
+ }
963
+ },
964
+ {
965
+ "name": "ort_oob_upper_bound_stream1",
966
+ "provenance": {
967
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
968
+ "test": "PositionIdsOOBPassthroughAllStreams",
969
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
970
+ },
971
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
972
+ "inputs": {
973
+ "x": {
974
+ "dtype": "float32",
975
+ "shape": [1, 1, 1, 6],
976
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
977
+ },
978
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [0, 2, 0] } },
979
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
980
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
981
+ },
982
+ "outputs": {
983
+ "y": {
984
+ "dtype": "float32",
985
+ "shape": [1, 1, 1, 6],
986
+ "data": { "kind": "values", "values": [-4.0, 2.0, -6.0, 1.0, 5.0, 3.0] }
987
+ }
988
+ }
989
+ },
990
+ {
991
+ "name": "ort_oob_saturated_high_stream1",
992
+ "provenance": {
993
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
994
+ "test": "PositionIdsOOBPassthroughAllStreams",
995
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
996
+ },
997
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
998
+ "inputs": {
999
+ "x": {
1000
+ "dtype": "float32",
1001
+ "shape": [1, 1, 1, 6],
1002
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
1003
+ },
1004
+ "positionIds": {
1005
+ "dtype": "int32",
1006
+ "shape": [3, 1, 1],
1007
+ "data": { "kind": "values", "values": [0, 2147483647, 0] }
1008
+ },
1009
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
1010
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
1011
+ },
1012
+ "outputs": {
1013
+ "y": {
1014
+ "dtype": "float32",
1015
+ "shape": [1, 1, 1, 6],
1016
+ "data": { "kind": "values", "values": [-4.0, 2.0, -6.0, 1.0, 5.0, 3.0] }
1017
+ }
1018
+ }
1019
+ },
1020
+ {
1021
+ "name": "ort_oob_saturated_low_stream1",
1022
+ "provenance": {
1023
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
1024
+ "test": "PositionIdsOOBPassthroughAllStreams",
1025
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
1026
+ },
1027
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
1028
+ "inputs": {
1029
+ "x": {
1030
+ "dtype": "float32",
1031
+ "shape": [1, 1, 1, 6],
1032
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
1033
+ },
1034
+ "positionIds": {
1035
+ "dtype": "int32",
1036
+ "shape": [3, 1, 1],
1037
+ "data": { "kind": "values", "values": [0, -2147483648, 0] }
1038
+ },
1039
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
1040
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
1041
+ },
1042
+ "outputs": {
1043
+ "y": {
1044
+ "dtype": "float32",
1045
+ "shape": [1, 1, 1, 6],
1046
+ "data": { "kind": "values", "values": [-4.0, 2.0, -6.0, 1.0, 5.0, 3.0] }
1047
+ }
1048
+ }
1049
+ },
1050
+ {
1051
+ "name": "ort_oob_negative_stream2",
1052
+ "provenance": {
1053
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
1054
+ "test": "PositionIdsOOBPassthroughAllStreams",
1055
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
1056
+ },
1057
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
1058
+ "inputs": {
1059
+ "x": {
1060
+ "dtype": "float32",
1061
+ "shape": [1, 1, 1, 6],
1062
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
1063
+ },
1064
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [0, 0, -1] } },
1065
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
1066
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
1067
+ },
1068
+ "outputs": {
1069
+ "y": {
1070
+ "dtype": "float32",
1071
+ "shape": [1, 1, 1, 6],
1072
+ "data": { "kind": "values", "values": [-4.0, -5.0, 3.0, 1.0, 2.0, 6.0] }
1073
+ }
1074
+ }
1075
+ },
1076
+ {
1077
+ "name": "ort_oob_upper_bound_stream2",
1078
+ "provenance": {
1079
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
1080
+ "test": "PositionIdsOOBPassthroughAllStreams",
1081
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
1082
+ },
1083
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
1084
+ "inputs": {
1085
+ "x": {
1086
+ "dtype": "float32",
1087
+ "shape": [1, 1, 1, 6],
1088
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
1089
+ },
1090
+ "positionIds": { "dtype": "int32", "shape": [3, 1, 1], "data": { "kind": "values", "values": [0, 0, 2] } },
1091
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
1092
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
1093
+ },
1094
+ "outputs": {
1095
+ "y": {
1096
+ "dtype": "float32",
1097
+ "shape": [1, 1, 1, 6],
1098
+ "data": { "kind": "values", "values": [-4.0, -5.0, 3.0, 1.0, 2.0, 6.0] }
1099
+ }
1100
+ }
1101
+ },
1102
+ {
1103
+ "name": "ort_oob_saturated_high_stream2",
1104
+ "provenance": {
1105
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
1106
+ "test": "PositionIdsOOBPassthroughAllStreams",
1107
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
1108
+ },
1109
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
1110
+ "inputs": {
1111
+ "x": {
1112
+ "dtype": "float32",
1113
+ "shape": [1, 1, 1, 6],
1114
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
1115
+ },
1116
+ "positionIds": {
1117
+ "dtype": "int32",
1118
+ "shape": [3, 1, 1],
1119
+ "data": { "kind": "values", "values": [0, 0, 2147483647] }
1120
+ },
1121
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
1122
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
1123
+ },
1124
+ "outputs": {
1125
+ "y": {
1126
+ "dtype": "float32",
1127
+ "shape": [1, 1, 1, 6],
1128
+ "data": { "kind": "values", "values": [-4.0, -5.0, 3.0, 1.0, 2.0, 6.0] }
1129
+ }
1130
+ }
1131
+ },
1132
+ {
1133
+ "name": "ort_oob_saturated_low_stream2",
1134
+ "provenance": {
1135
+ "source": "onnxruntime/test/contrib_ops/mrotary_embedding_op_test.cc",
1136
+ "test": "PositionIdsOOBPassthroughAllStreams",
1137
+ "notes": "Each invalid stream copies only its own rotation pair. Values outside int32 are saturated by the host projection."
1138
+ },
1139
+ "attrs": { "num_heads": 1, "rotary_embedding_dim": 6, "mrope_section": [1, 1, 1] },
1140
+ "inputs": {
1141
+ "x": {
1142
+ "dtype": "float32",
1143
+ "shape": [1, 1, 1, 6],
1144
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
1145
+ },
1146
+ "positionIds": {
1147
+ "dtype": "int32",
1148
+ "shape": [3, 1, 1],
1149
+ "data": { "kind": "values", "values": [0, 0, -2147483648] }
1150
+ },
1151
+ "cos": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 0.0 } },
1152
+ "sin": { "dtype": "float32", "shape": [2, 3], "data": { "kind": "constant", "value": 1.0 } }
1153
+ },
1154
+ "outputs": {
1155
+ "y": {
1156
+ "dtype": "float32",
1157
+ "shape": [1, 1, 1, 6],
1158
+ "data": { "kind": "values", "values": [-4.0, -5.0, 3.0, 1.0, 2.0, 6.0] }
1159
+ }
1160
+ }
1161
+ },
1162
+ {
1163
+ "name": "oob_mixed_rank3_float32_interleaved0",
1164
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 0, "interleaved": 0 },
1165
+ "inputs": {
1166
+ "x": { "dtype": "float32", "shape": [2, 2, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1167
+ "positionIds": {
1168
+ "dtype": "int32",
1169
+ "shape": [3, 2, 2],
1170
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1171
+ },
1172
+ "cos": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1173
+ "sin": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1174
+ },
1175
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 16], "tolerance": 0.00001, "relTolerance": 0.00001 } }
1176
+ },
1177
+ {
1178
+ "name": "oob_mixed_rank3_float32_interleaved1",
1179
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 1, "interleaved": 1 },
1180
+ "inputs": {
1181
+ "x": { "dtype": "float32", "shape": [2, 2, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1182
+ "positionIds": {
1183
+ "dtype": "int32",
1184
+ "shape": [3, 2, 2],
1185
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1186
+ },
1187
+ "cos": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1188
+ "sin": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1189
+ },
1190
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 16], "tolerance": 0.00001, "relTolerance": 0.00001 } }
1191
+ },
1192
+ {
1193
+ "name": "oob_mixed_rank3_float16_interleaved0",
1194
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 0, "interleaved": 0 },
1195
+ "inputs": {
1196
+ "x": { "dtype": "float16", "shape": [2, 2, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1197
+ "positionIds": {
1198
+ "dtype": "int32",
1199
+ "shape": [3, 2, 2],
1200
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1201
+ },
1202
+ "cos": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1203
+ "sin": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1204
+ },
1205
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 2, 16], "tolerance": 0.002, "relTolerance": 0.001 } }
1206
+ },
1207
+ {
1208
+ "name": "oob_mixed_rank3_float16_interleaved1",
1209
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 1, "interleaved": 1 },
1210
+ "inputs": {
1211
+ "x": { "dtype": "float16", "shape": [2, 2, 16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1212
+ "positionIds": {
1213
+ "dtype": "int32",
1214
+ "shape": [3, 2, 2],
1215
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1216
+ },
1217
+ "cos": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1218
+ "sin": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1219
+ },
1220
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 2, 16], "tolerance": 0.002, "relTolerance": 0.001 } }
1221
+ },
1222
+ {
1223
+ "name": "oob_mixed_rank4_float32_interleaved0",
1224
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 0, "interleaved": 0 },
1225
+ "inputs": {
1226
+ "x": { "dtype": "float32", "shape": [2, 2, 2, 8], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1227
+ "positionIds": {
1228
+ "dtype": "int32",
1229
+ "shape": [3, 2, 2],
1230
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1231
+ },
1232
+ "cos": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1233
+ "sin": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1234
+ },
1235
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 2, 8], "tolerance": 0.00001, "relTolerance": 0.00001 } }
1236
+ },
1237
+ {
1238
+ "name": "oob_mixed_rank4_float32_interleaved1",
1239
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 1, "interleaved": 1 },
1240
+ "inputs": {
1241
+ "x": { "dtype": "float32", "shape": [2, 2, 2, 8], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1242
+ "positionIds": {
1243
+ "dtype": "int32",
1244
+ "shape": [3, 2, 2],
1245
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1246
+ },
1247
+ "cos": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1248
+ "sin": { "dtype": "float32", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1249
+ },
1250
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 2, 2, 8], "tolerance": 0.00001, "relTolerance": 0.00001 } }
1251
+ },
1252
+ {
1253
+ "name": "oob_mixed_rank4_float16_interleaved0",
1254
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 0, "interleaved": 0 },
1255
+ "inputs": {
1256
+ "x": { "dtype": "float16", "shape": [2, 2, 2, 8], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1257
+ "positionIds": {
1258
+ "dtype": "int32",
1259
+ "shape": [3, 2, 2],
1260
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1261
+ },
1262
+ "cos": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1263
+ "sin": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1264
+ },
1265
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 2, 2, 8], "tolerance": 0.002, "relTolerance": 0.001 } }
1266
+ },
1267
+ {
1268
+ "name": "oob_mixed_rank4_float16_interleaved1",
1269
+ "attrs": { "num_heads": 2, "mrope_section": [2, 1, 1], "mrope_layout": 1, "interleaved": 1 },
1270
+ "inputs": {
1271
+ "x": { "dtype": "float16", "shape": [2, 2, 2, 8], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
1272
+ "positionIds": {
1273
+ "dtype": "int32",
1274
+ "shape": [3, 2, 2],
1275
+ "data": { "kind": "values", "values": [-1, 0, 1, 2, 0, -1, 2, 1, 2, 1, 0, -1] }
1276
+ },
1277
+ "cos": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": 0.1, "end": 0.8 } },
1278
+ "sin": { "dtype": "float16", "shape": [2, 4], "data": { "kind": "linspace", "start": -0.8, "end": -0.1 } }
1279
+ },
1280
+ "outputs": { "y": { "dtype": "float16", "shape": [2, 2, 2, 8], "tolerance": 0.002, "relTolerance": 0.001 } }
1281
  }
1282
  ]
1283
  }