use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use super::{check_arity, to_dense_f32, to_dense_i64, write_dense_f32};
pub struct RotaryEmbeddingKernel {
interleaved: bool,
num_heads: usize,
rotary_embedding_dim: usize,
}
pub struct RotaryEmbeddingFactory;
impl KernelFactory for RotaryEmbeddingFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let interleaved = node
.attr("interleaved")
.and_then(|a| a.as_int())
.unwrap_or(0)
!= 0;
let num_heads = node
.attr("num_heads")
.and_then(|a| a.as_int())
.unwrap_or(0)
.max(0) as usize;
let rotary_embedding_dim = node
.attr("rotary_embedding_dim")
.and_then(|a| a.as_int())
.unwrap_or(0)
.max(0) as usize;
Ok(Box::new(RotaryEmbeddingKernel {
interleaved,
num_heads,
rotary_embedding_dim,
}))
}
}
impl Kernel for RotaryEmbeddingKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("RotaryEmbedding", inputs, outputs, 3, 4, 1)?;
let x = to_dense_f32(&inputs[0])?;
let cos_cache = to_dense_f32(&inputs[1])?;
let sin_cache = to_dense_f32(&inputs[2])?;
let position_ids = if inputs.len() == 4 {
Some(to_dense_i64(&inputs[3])?)
} else {
None
};
let x_shape = inputs[0].shape;
let (batch, seq, heads, head_size, is_4d) = match x_shape.len() {
4 => {
(x_shape[0], x_shape[2], x_shape[1], x_shape[3], true)
}
3 => {
if self.num_heads == 0 {
return Err(EpError::KernelFailed(
"RotaryEmbedding: num_heads must be set for a 3D input".into(),
));
}
let hidden = x_shape[2];
if !hidden.is_multiple_of(self.num_heads) {
return Err(EpError::KernelFailed(format!(
"RotaryEmbedding: hidden {hidden} not divisible by num_heads {}",
self.num_heads
)));
}
(
x_shape[0],
x_shape[1],
self.num_heads,
hidden / self.num_heads,
false,
)
}
r => {
return Err(EpError::KernelFailed(format!(
"RotaryEmbedding: X must be rank 3 or 4, got rank {r}"
)));
}
};
let rotary_dim = if self.rotary_embedding_dim == 0 {
head_size
} else {
self.rotary_embedding_dim
};
if rotary_dim > head_size || !rotary_dim.is_multiple_of(2) {
return Err(EpError::KernelFailed(format!(
"RotaryEmbedding: rotary_embedding_dim {rotary_dim} invalid for head_size {head_size}"
)));
}
let half = rotary_dim / 2;
if x.is_empty() {
return write_dense_f32(&mut outputs[0], &[]);
}
if let Some(pos) = &position_ids {
let pos_shape = inputs[3].shape;
let expected = batch * seq;
if pos.len() != expected {
return Err(EpError::KernelFailed(format!(
"RotaryEmbedding: position_ids has {} elements, expected {expected} ([batch={batch}, seq={seq}]); shape {pos_shape:?}",
pos.len()
)));
}
}
let cache_stride = half; let cache_row = |b: usize, s: usize| -> Result<usize> {
let row = if let Some(pos) = &position_ids {
let p = pos[b * seq + s];
if p < 0 {
return Err(EpError::KernelFailed(
"RotaryEmbedding: negative position id".into(),
));
}
usize::try_from(p).map_err(|_| {
EpError::KernelFailed(
"RotaryEmbedding: position id exceeds supported range".into(),
)
})?
} else {
b * seq + s
};
let offset = row.checked_mul(cache_stride).ok_or_else(|| {
EpError::KernelFailed(format!(
"RotaryEmbedding: position {row} exceeds cos/sin cache extent"
))
})?;
let end = offset.checked_add(half).ok_or_else(|| {
EpError::KernelFailed(format!(
"RotaryEmbedding: position {row} exceeds cos/sin cache extent"
))
})?;
if offset > cos_cache.len()
|| end > cos_cache.len()
|| offset > sin_cache.len()
|| end > sin_cache.len()
{
return Err(EpError::KernelFailed(format!(
"RotaryEmbedding: position {row} exceeds cos/sin cache extent (row width {half})"
)));
}
Ok(offset)
};
let idx = |b: usize, h: usize, s: usize, d: usize| -> usize {
if is_4d {
((b * heads + h) * seq + s) * head_size + d
} else {
(b * seq + s) * (heads * head_size) + h * head_size + d
}
};
let mut y = vec![0.0f32; x.len()];
for b in 0..batch {
for s in 0..seq {
let crow = cache_row(b, s)?;
for h in 0..heads {
for k in 0..half {
let cos = cos_cache[crow + k];
let sin = sin_cache[crow + k];
let (d1, d2) = if self.interleaved {
(2 * k, 2 * k + 1)
} else {
(k, k + half)
};
let x1 = x[idx(b, h, s, d1)];
let x2 = x[idx(b, h, s, d2)];
y[idx(b, h, s, d1)] = cos * x1 - sin * x2;
y[idx(b, h, s, d2)] = sin * x1 + cos * x2;
}
for d in rotary_dim..head_size {
y[idx(b, h, s, d)] = x[idx(b, h, s, d)];
}
}
}
}
write_dense_f32(&mut outputs[0], &y)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
#[test]
fn rope_half_rotation_hand_computed() {
let c0 = 0.5f32;
let c1 = 0.8f32;
let s0 = (1.0f32 - c0 * c0).sqrt();
let s1 = (1.0f32 - c1 * c1).sqrt();
let x = Owned::f32(&[1, 1, 1, 4], &[1., 2., 3., 4.]);
let cos = Owned::f32(&[1, 1, 2], &[c0, c1]);
let sin = Owned::f32(&[1, 1, 2], &[s0, s1]);
let mut out = Owned::zeros_f32(&[1, 1, 1, 4]);
RotaryEmbeddingKernel {
interleaved: false,
num_heads: 0,
rotary_embedding_dim: 0,
}
.execute(&[x.view(), cos.view(), sin.view()], &mut [out.view_mut()])
.unwrap();
let want = [
c0 * 1.0 - s0 * 3.0,
c1 * 2.0 - s1 * 4.0,
s0 * 1.0 + c0 * 3.0,
s1 * 2.0 + c1 * 4.0,
];
for (g, w) in out.to_f32().iter().zip(&want) {
assert!((g - w).abs() < 1e-6, "got {g}, want {w}");
}
}
#[test]
fn rope_interleaved_hand_computed() {
let c0 = 0.5f32;
let c1 = 0.8f32;
let s0 = (1.0f32 - c0 * c0).sqrt();
let s1 = (1.0f32 - c1 * c1).sqrt();
let x = Owned::f32(&[1, 1, 1, 4], &[1., 2., 3., 4.]);
let cos = Owned::f32(&[1, 1, 2], &[c0, c1]);
let sin = Owned::f32(&[1, 1, 2], &[s0, s1]);
let mut out = Owned::zeros_f32(&[1, 1, 1, 4]);
RotaryEmbeddingKernel {
interleaved: true,
num_heads: 0,
rotary_embedding_dim: 0,
}
.execute(&[x.view(), cos.view(), sin.view()], &mut [out.view_mut()])
.unwrap();
let want = [
c0 * 1.0 - s0 * 2.0,
s0 * 1.0 + c0 * 2.0,
c1 * 3.0 - s1 * 4.0,
s1 * 3.0 + c1 * 4.0,
];
for (g, w) in out.to_f32().iter().zip(&want) {
assert!((g - w).abs() < 1e-6, "got {g}, want {w}");
}
}
#[test]
fn rope_zero_angle_is_identity() {
let x = Owned::f32(&[1, 2, 1, 4], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let cos = Owned::f32(&[1, 1, 2], &[1., 1.]);
let sin = Owned::f32(&[1, 1, 2], &[0., 0.]);
let mut out = Owned::zeros_f32(&[1, 2, 1, 4]);
RotaryEmbeddingKernel {
interleaved: false,
num_heads: 0,
rotary_embedding_dim: 0,
}
.execute(&[x.view(), cos.view(), sin.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![1., 2., 3., 4., 5., 6., 7., 8.]);
}
#[test]
fn rope_3d_with_num_heads_and_position_ids() {
let x = Owned::f32(&[1, 2, 4], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let cos = Owned::f32(&[2, 1], &[1.0, 0.0]); let sin = Owned::f32(&[2, 1], &[0.0, 1.0]); let pos = Owned::i64(&[1, 2], &[0, 1]);
let mut out = Owned::zeros_f32(&[1, 2, 4]);
RotaryEmbeddingKernel {
interleaved: false,
num_heads: 2,
rotary_embedding_dim: 0,
}
.execute(
&[x.view(), cos.view(), sin.view(), pos.view()],
&mut [out.view_mut()],
)
.unwrap();
let want = [1., 2., 3., 4., -6., 5., -8., 7.];
for (g, w) in out.to_f32().iter().zip(&want) {
assert!((g - w).abs() < 1e-6, "got {g}, want {w}");
}
}
#[test]
fn rope_partial_rotary_dim_passes_through_tail() {
let x = Owned::f32(&[1, 1, 1, 4], &[1., 2., 3., 4.]);
let cos = Owned::f32(&[1, 1, 1], &[0.0]);
let sin = Owned::f32(&[1, 1, 1], &[1.0]);
let mut out = Owned::zeros_f32(&[1, 1, 1, 4]);
RotaryEmbeddingKernel {
interleaved: false,
num_heads: 0,
rotary_embedding_dim: 2,
}
.execute(&[x.view(), cos.view(), sin.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![-2., 1., 3., 4.]);
}
#[test]
fn rope_zero_sized_input_returns_empty() {
let x = Owned::f32(&[0, 1, 1, 4], &[]);
let cos = Owned::f32(&[1, 1, 2], &[1., 1.]);
let sin = Owned::f32(&[1, 1, 2], &[0., 0.]);
let mut out = Owned::zeros_f32(&[0, 1, 1, 4]);
RotaryEmbeddingKernel {
interleaved: false,
num_heads: 0,
rotary_embedding_dim: 0,
}
.execute(&[x.view(), cos.view(), sin.view()], &mut [out.view_mut()])
.unwrap();
assert!(out.to_f32().is_empty());
}
#[test]
fn rope_out_of_range_position_errors() {
let x = Owned::f32(&[1, 2, 4], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let cos = Owned::f32(&[2, 1], &[1.0, 0.0]);
let sin = Owned::f32(&[2, 1], &[0.0, 1.0]);
let pos = Owned::i64(&[1, 2], &[0, 5]);
let mut out = Owned::zeros_f32(&[1, 2, 4]);
let err = RotaryEmbeddingKernel {
interleaved: false,
num_heads: 2,
rotary_embedding_dim: 0,
}
.execute(
&[x.view(), cos.view(), sin.view(), pos.view()],
&mut [out.view_mut()],
);
assert!(err.is_err(), "out-of-range position must return an error");
}
#[test]
fn rope_i64_max_position_errors_without_overflow() {
let x = Owned::f32(&[1, 1, 4], &[1., 2., 3., 4.]);
let cos = Owned::f32(&[1, 2], &[1.0, 1.0]);
let sin = Owned::f32(&[1, 2], &[0.0, 0.0]);
let pos = Owned::i64(&[1, 1], &[i64::MAX]);
let mut out = Owned::zeros_f32(&[1, 1, 4]);
let err = RotaryEmbeddingKernel {
interleaved: false,
num_heads: 2,
rotary_embedding_dim: 0,
}
.execute(
&[x.view(), cos.view(), sin.view(), pos.view()],
&mut [out.view_mut()],
);
assert!(err.is_err(), "i64::MAX position must return an error");
}
#[test]
fn rope_negative_position_errors() {
let x = Owned::f32(&[1, 1, 4], &[1., 2., 3., 4.]);
let cos = Owned::f32(&[1, 2], &[1.0, 1.0]);
let sin = Owned::f32(&[1, 2], &[0.0, 0.0]);
let pos = Owned::i64(&[1, 1], &[-1]);
let mut out = Owned::zeros_f32(&[1, 1, 4]);
let err = RotaryEmbeddingKernel {
interleaved: false,
num_heads: 2,
rotary_embedding_dim: 0,
}
.execute(
&[x.view(), cos.view(), sin.view(), pos.view()],
&mut [out.view_mut()],
);
assert!(err.is_err(), "negative position must return an error");
}
#[test]
fn rope_bad_position_ids_shape_errors() {
let x = Owned::f32(&[1, 2, 4], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let cos = Owned::f32(&[2, 1], &[1.0, 0.0]);
let sin = Owned::f32(&[2, 1], &[0.0, 1.0]);
let pos = Owned::i64(&[1, 1], &[0]);
let mut out = Owned::zeros_f32(&[1, 2, 4]);
let err = RotaryEmbeddingKernel {
interleaved: false,
num_heads: 2,
rotary_embedding_dim: 0,
}
.execute(
&[x.view(), cos.view(), sin.view(), pos.view()],
&mut [out.view_mut()],
);
assert!(err.is_err(), "malformed position_ids must return an error");
}
}