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

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

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.

Downloads last month
-
apache-2.0
Supported hardwares new
CUDA
8.08.68.99.010.012.0
GPU
B300
288GB
NVIDIA SXM
B200
192GB
NVIDIA SXM
H200
141GB
NVIDIA SXM
H100
80GB
GPU
H800
80GB
GPU
H20
96GB
GPU
L40s
48GB
GPU
L40
48GB
GPU
L20
48GB
GPU
L4
24GB
DGX Spark
GB10
128GB
GPU
RTX PRO 6000 WS
96GB
GPU
RTX PRO 6000 Max-Q
96GB
GPU
RTX PRO 5000
48GB
GPU
RTX PRO 4500 WS
32GB
GPU
RTX PRO 4000
24GB
GPU
RTX PRO 4000 SFF
24GB
GPU
RTX PRO 2000
16GB
GPU
RTX 6000 Ada
48GB
GPU
RTX 5880 Ada
48GB
RTX
RTX 5000 Ada
32GB
GPU
RTX 4500 Ada
24GB
RTX
RTX 4000 Ada
20GB
RTX
RTX 4000 SFF Ada
20GB
GPU
RTX 3500 Ada Mobile
12GB
GPU
RTX 2000 Ada
16GB
GPU
RTX A6000
48GB
GPU
RTX A5000
8GB
GPU
RTX A5000 Max-Q
16GB
GPU
RTX A5000 Mobile
16GB
GPU
RTX A4000
16GB
GPU
RTX A4000 Max-Q
8GB
GPU
RTX A4000 Mobile
8GB
GPU
RTX A3000 Mobile
6GB
GPU
RTX A2000
6GB
GPU
RTX A2000 Embedded
4GB
GPU
RTX A2000 Max-Q
4GB
GPU
RTX A2000 Mobile
4GB
GPU
A800
40GB
GPU
A100
80GB
GPU
A40
48GB
GPU
A30
24GB
GPU
A10
24GB
GPU
A2
16GB
RTX
RTX 5090
32GB
RTX
RTX 5090 D
32GB
RTX
RTX 5090 Mobile
24GB
RTX
RTX 5080
16GB
RTX
RTX 5080 Mobile
16GB
RTX
RTX 5070
12GB
RTX
RTX 5070 Mobile
8GB
RTX
RTX 5070 Ti
16GB
RTX
RTX 5070 Ti Mobile
12GB
RTX
RTX 5060 Ti
16GB
RTX
RTX 5060
8GB
RTX
RTX 5060 Mobile
8GB
RTX
RTX 5050
8GB
RTX
RTX 5050 Mobile
8GB
RTX
RTX 4090
24GB
RTX
RTX 4090D
24GB
RTX
RTX 4090 Mobile
16GB
RTX
RTX 4080 SUPER
16GB
RTX
RTX 4080
16GB
RTX
RTX 4080 Mobile
12GB
RTX
RTX 4070
12GB
RTX
RTX 4070 Mobile
8GB
RTX
RTX 4070 Ti
12GB
RTX
RTX 4070 Super
12GB
RTX
RTX 4070 Ti Super
16GB
RTX
RTX 4060
8GB
RTX
RTX 4060 Ti
8GB
RTX
RTX 4090 Laptop
16GB
RTX
RTX 4080 Laptop
12GB
RTX
RTX 4070 Laptop
8GB
RTX
RTX 4060 Laptop
8GB
RTX
RTX 4050 Laptop
6GB
RTX
RTX 3090
24GB
RTX
RTX 3090 Ti
24GB
RTX
RTX 3080
12GB
RTX
RTX 3080 Ti
12GB
RTX
RTX 3080 Mobile
16GB
RTX
RTX 3070
8GB
RTX
RTX 3070 Ti
8GB
RTX
RTX 3070 Ti Mobile
8GB
RTX
RTX 3060 Ti
8GB
RTX
RTX 3060
12GB
RTX
RTX 3060 Mobile
6GB
RTX
RTX 3050 Mobile
4GB
GPU
RTX 2050 Mobile
4GB
Jetson
Jetson AGX Orin 64GB
64GB
Jetson
Jetson AGX Orin 32GB
32GB
Jetson
Jetson Orin NX 16GB
16GB
Jetson
Jetson Orin NX 8GB
8GB
Jetson
Jetson Orin Nano 8GB
8GB
Jetson
Jetson Orin Nano 4GB
4GB
OS
linux
Arch
x86_64
Kernel Builder
19aaa64