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