File size: 6,774 Bytes
2437bce
4eccb3a
2437bce
4eccb3a
 
 
 
2437bce
4eccb3a
 
 
 
 
 
1844ba1
4eccb3a
 
 
 
 
1844ba1
4eccb3a
1844ba1
 
 
 
4eccb3a
 
 
1844ba1
4eccb3a
1844ba1
4eccb3a
 
 
 
 
 
 
 
 
1844ba1
4eccb3a
 
 
 
 
 
 
 
6332f1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4eccb3a
 
1844ba1
4eccb3a
 
6332f1b
 
1844ba1
4eccb3a
 
 
 
1844ba1
6332f1b
1844ba1
 
 
4eccb3a
 
1844ba1
4eccb3a
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
---
library_name: kernels
license: apache-2.0
tags:
- kernel
- webgpu
- wgsl
---
# ai.onnx.RotaryEmbedding

`ai.onnx`  ·  standard ONNX operator  ·  ONNX opset ≥ 23

## Description

Implements ONNX opset-23 RotaryEmbedding for float16 and float32 tensors. Applies rotary positional embeddings (RoPE) by rotating each head's embedding vector using precomputed `cos_cache` and `sin_cache` values. A partial rotation can be applied by setting `rotary_embedding_dim` to rotate only a prefix of the head dimension. `position_ids` keeps its standard logical int64 type; valid positions are non-negative and bounded by the WebGPU-addressable cache, so the backend stores them losslessly as uint32. Other ONNX floating-point input types are unsupported.

See the [ONNX `RotaryEmbedding` spec](https://onnx.ai/onnx/operators/onnx__RotaryEmbedding.html) for the reference semantics.

## Inputs

| Name | Upstream name | Logical dtype | WebGPU storage | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- | --- | --- |
| `x` | `X` | `T` | same as logical dtype | — | — | Input token embeddings. Shape is `(batch_size, sequence_length, hidden_size)` for rank 3 or `(batch_size, num_heads, sequence_length, head_size)` for rank 4. `head_size` must be even, and the `num_heads` attribute is required for rank-3 input. | required |
| `cos` | `cos_cache` | `T` | same as logical dtype | — | — | Precomputed cosine values. Without `position_ids`, shape is `(batch_size, sequence_length, rotary_dim/2)`; with `position_ids`, shape is `(max_sequence_length, rotary_dim/2)`. | required |
| `sin` | `sin_cache` | `T` | same as logical dtype | — | — | Precomputed sine values with the same shape and type as `cos_cache`. | required |
| `positionIds` | `position_ids` | `M` | `uint32` | `2` | — | Optional logical int64 per-token position indices of shape `(batch_size, sequence_length)`. Valid positions are non-negative cache-row indices and use uint32 WebGPU storage. When supplied, the 2D cache tables are gathered at these positions. | optional |

## Outputs

| Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- | --- |
| `y` | `Y` | `T` | same as `x` | same as `x` | Rotary-position-encoded tensor with the same shape and type as `X`. | required |

## Attributes

Attributes and default values (overridable per request):

| Attribute | Default | Description |
| --- | --- | --- |
| `interleaved` | `0` | Set to 1 to rotate using an interleaved pattern (even/odd elements), or 0 to split the head dimension into two contiguous halves. Default is 0. |
| `num_heads` | — | Optional number of attention heads. ONNX requires this attribute when `X` is rank 3; it is unnecessary for rank-4 input because the head count is explicit in the shape. |
| `rotary_embedding_dim` | `0` | Number of head-dimension elements to rotate; `0` means rotate the full head dimension. When set, only the leading `rotary_embedding_dim` elements are rotated and the rest are passed through unchanged. |

## Type constraints

| Variable | Allowed dtypes |
| --- | --- |
| `T` | `float32`, `float16` |
| `M` | `int64` |

## Implementation variants

One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.

- `rank3_cache2_pos_quad` — Processes four pairs with vector loads and stores when head and cache alignment keep each group within a head and rotation region.
- `rank3_cache2_pos` — Processes one rotated pair per invocation with scalar storage, or copies an unchanged pair from the tail.
- `rank3_cache2_pos_output_vec4` — Writes four contiguous outputs per invocation with gathered pair values. Non-rotated components retain their original storage bits.
- `rank4_cache2_pos_quad` — Processes four pairs with vector loads and stores when head and cache alignment keep each group within a head and rotation region.
- `rank4_cache2_pos` — Processes one rotated pair per invocation with scalar storage, or copies an unchanged pair from the tail.
- `rank4_cache2_pos_output_vec4` — Writes four contiguous outputs per invocation with gathered pair values. Non-rotated components retain their original storage bits.
- `rank3_cache3_quad` — Processes four pairs with vector loads and stores when head and cache alignment keep each group within a head and rotation region.
- `rank3_cache3` — Processes one rotated pair per invocation with scalar storage, or copies an unchanged pair from the tail.
- `rank3_cache3_output_vec4` — Writes four contiguous outputs per invocation with gathered pair values. Non-rotated components retain their original storage bits.
- `rank4_cache3_quad` — Processes four pairs with vector loads and stores when head and cache alignment keep each group within a head and rotation region.
- `rank4_cache3` — Processes one rotated pair per invocation with scalar storage, or copies an unchanged pair from the tail.
- `rank4_cache3_output_vec4` — Writes four contiguous outputs per invocation with gathered pair values. Non-rotated components retain their original storage bits.

## Files

- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
- [`test.json`](build/webgpu/test.json) — correctness cases
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
- [`rotary-embedding-output-vec4.wgsl.jinja`](build/webgpu/rotary-embedding-output-vec4.wgsl.jinja)
- [`rotary-embedding-quad.wgsl.jinja`](build/webgpu/rotary-embedding-quad.wgsl.jinja)
- [`rotary-embedding.wgsl.jinja`](build/webgpu/rotary-embedding.wgsl.jinja)

## Use with `@huggingface/kernels`

```sh
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
```

Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.

The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
It follows the `v1` branch as fixes land. To pin exact artifact bytes, pass a 40-character commit `revision` instead of `version`.

Replace each `*Data` placeholder with a typed array containing the corresponding input data.

```js
import { getKernel } from "@huggingface/kernels";

const kernel = await getKernel("webgpu-kernels/ai.onnx.RotaryEmbedding", { version: 1 });
const { y } = await kernel({
  x: { data: xData, shape: [1, 2, 1, 4] },
  cos: { data: cosData, shape: [16, 2] },
  sin: { data: sinData, shape: [16, 2] },
  positionIds: { data: positionIdsData, shape: [1, 1] },
});
```