use crate::error::{WhisperError, WhisperResult};
#[derive(Debug, Clone)]
pub struct RopeConfig {
pub head_dim: usize,
pub base: f32,
pub max_seq_len: usize,
pub rotary_dim: Option<usize>,
}
impl RopeConfig {
#[must_use]
pub fn lfm2_2_6b() -> Self {
Self {
head_dim: 64,
base: 1_000_000.0, max_seq_len: 4096, rotary_dim: None, }
}
#[must_use]
pub fn effective_rotary_dim(&self) -> usize {
self.rotary_dim.unwrap_or(self.head_dim)
}
pub fn validate(&self) -> WhisperResult<()> {
let rdim = self.effective_rotary_dim();
if rdim % 2 != 0 {
return Err(WhisperError::Model(format!(
"rotary_dim ({rdim}) must be even for RoPE rotation pairs",
)));
}
if rdim > self.head_dim {
return Err(WhisperError::Model(format!(
"rotary_dim ({rdim}) must be <= head_dim ({})",
self.head_dim
)));
}
if self.base <= 0.0 {
return Err(WhisperError::Model("base must be positive".into()));
}
if self.max_seq_len == 0 {
return Err(WhisperError::Model("max_seq_len must be > 0".into()));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct RotaryEmbedding {
pub config: RopeConfig,
cos_cache: Vec<f32>,
sin_cache: Vec<f32>,
}
impl RotaryEmbedding {
pub fn new(config: RopeConfig) -> WhisperResult<Self> {
config.validate()?;
let rotary_dim = config.effective_rotary_dim();
let half_rotary = rotary_dim / 2;
let max_seq = config.max_seq_len;
let inv_freq: Vec<f32> = (0..half_rotary)
.map(|i| {
let exp = -2.0 * (i as f32) / (rotary_dim as f32);
config.base.powf(exp)
})
.collect();
let mut cos_cache = vec![0.0f32; max_seq * half_rotary];
let mut sin_cache = vec![0.0f32; max_seq * half_rotary];
for pos in 0..max_seq {
for (i, &freq) in inv_freq.iter().enumerate() {
let angle = (pos as f32) * freq;
cos_cache[pos * half_rotary + i] = angle.cos();
sin_cache[pos * half_rotary + i] = angle.sin();
}
}
Ok(Self {
config,
cos_cache,
sin_cache,
})
}
pub fn forward(
&self,
x: &[f32],
seq_len: usize,
num_heads: usize,
position_offset: usize,
) -> WhisperResult<Vec<f32>> {
let head_dim = self.config.head_dim;
let rotary_dim = self.config.effective_rotary_dim();
let half_rotary = rotary_dim / 2;
let max_seq = self.config.max_seq_len;
let expected_len = seq_len * num_heads * head_dim;
if x.len() != expected_len {
return Err(WhisperError::Model(format!(
"input length {} != expected {} (seq={}, heads={}, dim={})",
x.len(),
expected_len,
seq_len,
num_heads,
head_dim
)));
}
if position_offset + seq_len > max_seq {
return Err(WhisperError::Model(format!(
"position {} + seq_len {} exceeds max_seq_len {}",
position_offset, seq_len, max_seq
)));
}
let mut output = x.to_vec();
for s in 0..seq_len {
let pos = position_offset + s;
let cos = &self.cos_cache[pos * half_rotary..(pos + 1) * half_rotary];
let sin = &self.sin_cache[pos * half_rotary..(pos + 1) * half_rotary];
for h in 0..num_heads {
let offset = (s * num_heads + h) * head_dim;
for i in 0..half_rotary {
let x0 = x[offset + 2 * i];
let x1 = x[offset + 2 * i + 1];
output[offset + 2 * i] = x0 * cos[i] - x1 * sin[i];
output[offset + 2 * i + 1] = x0 * sin[i] + x1 * cos[i];
}
}
}
Ok(output)
}
pub fn forward_inplace(
&self,
x: &mut [f32],
seq_len: usize,
num_heads: usize,
position_offset: usize,
) -> WhisperResult<()> {
let head_dim = self.config.head_dim;
let rotary_dim = self.config.effective_rotary_dim();
let half_rotary = rotary_dim / 2;
let max_seq = self.config.max_seq_len;
if position_offset + seq_len > max_seq {
return Err(WhisperError::Model(format!(
"position {} + seq_len {} exceeds max_seq_len {}",
position_offset, seq_len, max_seq
)));
}
for s in 0..seq_len {
let pos = position_offset + s;
let cos = &self.cos_cache[pos * half_rotary..(pos + 1) * half_rotary];
let sin = &self.sin_cache[pos * half_rotary..(pos + 1) * half_rotary];
for h in 0..num_heads {
let offset = (s * num_heads + h) * head_dim;
for i in 0..half_rotary {
let x0 = x[offset + 2 * i];
let x1 = x[offset + 2 * i + 1];
x[offset + 2 * i] = x0 * cos[i] - x1 * sin[i];
x[offset + 2 * i + 1] = x0 * sin[i] + x1 * cos[i];
}
}
}
Ok(())
}
#[must_use]
pub fn get_cos(&self, start: usize, len: usize) -> &[f32] {
let half_rotary = self.config.effective_rotary_dim() / 2;
&self.cos_cache[start * half_rotary..(start + len) * half_rotary]
}
#[must_use]
pub fn get_sin(&self, start: usize, len: usize) -> &[f32] {
let half_rotary = self.config.effective_rotary_dim() / 2;
&self.sin_cache[start * half_rotary..(start + len) * half_rotary]
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
2 * self.cos_cache.len() * std::mem::size_of::<f32>()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rope_config_lfm2() {
let config = RopeConfig::lfm2_2_6b();
assert_eq!(config.head_dim, 64);
assert!((config.base - 1_000_000.0).abs() < 1.0);
assert_eq!(config.max_seq_len, 4096);
assert!(config.validate().is_ok());
}
#[test]
fn test_rope_config_validation() {
let config = RopeConfig {
head_dim: 63,
base: 10000.0,
max_seq_len: 100,
rotary_dim: None,
};
assert!(config.validate().is_err());
let config = RopeConfig {
head_dim: 64,
base: 0.0,
max_seq_len: 100,
rotary_dim: None,
};
assert!(config.validate().is_err());
}
#[test]
fn test_rope_new() {
let config = RopeConfig {
head_dim: 8,
base: 10000.0,
max_seq_len: 10,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
assert_eq!(rope.cos_cache.len(), 10 * 4); assert_eq!(rope.sin_cache.len(), 10 * 4);
}
#[test]
fn test_rope_forward_shape() {
let config = RopeConfig {
head_dim: 8,
base: 10000.0,
max_seq_len: 100,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let seq_len = 5;
let num_heads = 4;
let head_dim = 8;
let input = vec![1.0f32; seq_len * num_heads * head_dim];
let output = rope
.forward(&input, seq_len, num_heads, 0)
.expect("forward should succeed");
assert_eq!(output.len(), input.len());
}
#[test]
fn test_rope_rotation_preserves_norm() {
let config = RopeConfig {
head_dim: 4,
base: 10000.0,
max_seq_len: 10,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let input = vec![1.0, 2.0, 3.0, 4.0];
let input_norm: f32 = input.iter().map(|x| x * x).sum::<f32>().sqrt();
let output = rope
.forward(&input, 1, 1, 0)
.expect("forward should succeed");
let output_norm: f32 = output.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(input_norm - output_norm).abs() < 1e-5,
"Rotation should preserve norm: {} vs {}",
input_norm,
output_norm
);
}
#[test]
fn test_rope_position_0_is_identity() {
let config = RopeConfig {
head_dim: 4,
base: 10000.0,
max_seq_len: 10,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let input = vec![1.0, 2.0, 3.0, 4.0];
let output = rope
.forward(&input, 1, 1, 0)
.expect("forward should succeed");
for (i, &v) in output.iter().enumerate() {
assert!(
(v - input[i]).abs() < 1e-5,
"Position 0 should be identity: {} vs {}",
v,
input[i]
);
}
}
#[test]
fn test_rope_different_positions_differ() {
let config = RopeConfig {
head_dim: 4,
base: 10000.0,
max_seq_len: 10,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let input = vec![1.0, 2.0, 3.0, 4.0];
let output_pos0 = rope
.forward(&input, 1, 1, 0)
.expect("forward should succeed");
let output_pos5 = rope
.forward(&input, 1, 1, 5)
.expect("forward should succeed");
let diff: f32 = output_pos0
.iter()
.zip(output_pos5.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(diff > 0.01, "Different positions should differ");
}
#[test]
fn test_rope_inplace() {
let config = RopeConfig {
head_dim: 4,
base: 10000.0,
max_seq_len: 10,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let original = vec![1.0, 2.0, 3.0, 4.0];
let mut inplace = original.clone();
let output = rope
.forward(&original, 1, 1, 3)
.expect("forward should succeed");
rope.forward_inplace(&mut inplace, 1, 1, 3)
.expect("inplace should succeed");
for (i, &v) in inplace.iter().enumerate() {
assert!(
(v - output[i]).abs() < 1e-6,
"Inplace should match forward: {} vs {}",
v,
output[i]
);
}
}
#[test]
fn test_rope_position_overflow() {
let config = RopeConfig {
head_dim: 4,
base: 10000.0,
max_seq_len: 10,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let input = vec![1.0f32; 4];
let result = rope.forward(&input, 1, 1, 15);
assert!(result.is_err());
}
#[test]
fn test_rope_memory() {
let config = RopeConfig {
head_dim: 64,
base: 10000.0,
max_seq_len: 4096,
rotary_dim: None,
};
let rope = RotaryEmbedding::new(config).expect("should create RoPE");
let expected = 2 * 4096 * 32 * 4;
assert_eq!(rope.memory_bytes(), expected);
}
}