mrope / README.md
phanerozoic's picture
Card: standardized form with hero
6c1bf8d verified
|
Raw
History Blame
5.1 kB
---
library_name: kernels
license: apache-2.0
---
# mrope
Interleaved multimodal rotary position embedding, applied in place, loadable
through `kernels`. The reference baseline is transformers'
`Qwen3_5TextRotaryEmbedding`, matched to 3.1e-7 in fp32, with an fp64
rotation as the accuracy oracle.
A multimodal model needs three positions per token rather than one: an image
patch has a place in time, in height, and in width, so the rotary frequency
band is partitioned across those axes. The interleaved layout cycles them,
`[THWTHW...]`, so each axis keeps a full spread of wavelengths. This kernel
applies all of it in one pass, one block per token with the angles computed
once into shared memory and reused across every head, and evaluates angles in
double precision, which is what keeps a 262,144-token context accurate where
the eager fp32 path has lost four orders of magnitude.
![Thirty-two frequency dials interleaved by axis: advancing one axis's position spins only that axis's dials](https://huggingface.co/kernels/phanerozoic/mrope/resolve/main/media/hero.gif)
*The kernel's own axis map probed live: 32 frequency dials colored time,
height, width in the interleaved pattern, each phase advancing one axis and
spinning only that axis's dials. At position 262,143 the kernel's angle errs
9e-8 radians against an exact reference; the fp32 eager path errs 8e-5.*
## Usage
```python
import torch
from kernels import get_kernel
mrope = get_kernel("phanerozoic/mrope", version=1, trust_remote_code=True)
# q is [B, S, Hq, D], k is [B, S, Hk, D], pos is [3, B, S] int32 holding the
# time, height and width position of every token. Both rotate in place.
mrope.apply_mrope(q, k, pos, head_dim=256, partial_rotary_factor=0.25,
mrope_section=(11, 11, 10), rope_theta=1e7)
```
`version` selects the release branch; `trust_remote_code` is required by
`kernels` for publishers without the trusted-publisher mark.
## API
| Symbol | Purpose |
|---|---|
| `apply_mrope(q, k, pos, head_dim, partial_rotary_factor, mrope_section, rope_theta)` | rotate in place; `k` may be `None` |
| `MRope(head_dim, partial_rotary_factor, mrope_section, rope_theta)` | module form holding the geometry |
| `axis_of_frequency(half, section)` | which position axis each frequency uses |
## Method
Two details make this its own kernel rather than a case of ordinary RoPE.
Partial rotary: only the first `head_dim * partial_rotary_factor` channels
rotate, the rest pass through untouched; at head dimension 256 and factor
0.25 that is 64 rotated channels of 256. Interleaved sections: frequency `j`
uses the time position unless `j % 3 == 1 and j < 3 * section[1]` (height)
or `j % 3 == 2 and j < 3 * section[2]` (width); with sections `[11, 11, 10]`
over 32 frequencies that reduces exactly to `j % 3`. One block per
(batch, token) computes the 32 angles once into shared memory and reuses
them across every head of that token.
## Measured
Against the eager transformers implementation, head dimension 256, 24 query
and 4 key heads:
| batch | tokens | eager | this | speedup |
|---|---|---|---|---|
| 1 | 1 | 0.3107 ms | 0.0243 ms | 12.80x |
| 1 | 512 | 0.5334 ms | 0.0215 ms | 24.76x |
| 1 | 4,096 | 0.7179 ms | 0.1809 ms | 3.97x |
| 4 | 2,048 | 1.4754 ms | 0.4162 ms | 3.54x |
The eager path builds a `[3, B, S, 32]` frequency tensor, slices it three
ways, concatenates, takes a cosine and a sine, and runs four elementwise
passes; this kernel touches each element once.
Accuracy at long positions, against an independent fp64 rotation:
| position | this kernel | eager reference |
|---|---|---|
| 1,024 | 1.6e-8 | 2.1e-6 |
| 65,536 | 1.9e-8 | 1.4e-4 |
| 262,143 | 1.8e-8 | 7.2e-4 |
Angles are evaluated in double before being reduced to float; the eager
reference computes them in float32, which loses the fractional part of a
cycle at large positions.
## Correctness
Verified against transformers 5.13:
- The frequency-to-axis assignment is `torch.equal` to the reference slice
construction, and per-axis counts equal the configured sections.
- Output matches to 3.1e-7 relative in fp32 and 1.0e-3 in bf16, across batch
1 to 3, 8 to 512 tokens, grouped and ungrouped head counts, with distinct
positions on all three axes.
- Moving only the height position changes exactly the frequencies assigned
to height, checked elementwise; each rotation pair keeps its magnitude to
1e-3; the tail beyond `rotary_dim` is bitwise unchanged; position zero is
the exact identity.
## Requirements and limits
- NVIDIA GPU with compute capability 8.0+.
- `q` and `k` contiguous `[B, S, H, D]`, bf16 or fp32, sharing a dtype;
`[B, H, S, D]` tensors need a permute first.
- `rotary_dim / 2` at most 128; positions int32 (contexts to 2^31 tokens).
- Forward only: RoPE is its own inverse up to a sign, so backward is a
second call with negated positions.
## References
Su et al., "RoFormer" (2021); Wang et al., "Qwen2-VL" (2024) for multimodal
rotary sections; the interleaved layout as implemented in transformers'
`Qwen3_5TextRotaryEmbedding`.
## License
Apache-2.0.