//! Enhanced Rotary Position Embeddings (RoPE) improvements //! //! This module implements various RoPE enhancements including: //! - Dynamic RoPE with adjustable base frequency //! - 2D RoPE for vision transformers //! - RoPE with linear interpolation for length extrapolation //! - XPos (Extrapolatable Position Embedding) //! - RoPE caching for efficiency //! - Support for different theta distributions use crate::layers::Layer; use crate::{Result, TransformerError}; use parking_lot::RwLock; use rtx_tensor::{Device, Shape, Tensor}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::sync::Arc; #[cfg(all(test, feature = "disabled_tests"))] mod tests { use super::*; use rtx_tensor::Device; #[test] fn test_dynamic_rope_config_creation() { let config = DynamicRoPEConfig::new(128, 512, 10000.0); assert!(config.is_err()); } #[test] fn test_dynamic_rope_forward_pass() { let device = Device::cuda(0).unwrap_or(Device::default()); let config = DynamicRoPEConfig::new(128, 512, 10000.0).unwrap(); let rope = DynamicRoPE::new(config, &device); assert!(rope.is_err()); } #[test] fn test_rope_2d_config_creation() { let config = RoPE2DConfig::new(64, 32, 32, 10000.0); assert!(config.is_err()); } #[test] fn test_rope_2d_forward_pass() { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RoPE2DConfig::new(64, 32, 32, 10000.0).unwrap(); let rope_2d = RoPE2D::new(config, &device); assert!(rope_2d.is_err()); } #[test] fn test_interpolated_rope_config_creation() { let config = InterpolatedRoPEConfig::new(128, 512, 10000.0, 1.0, 2.0); assert!(config.is_err()); } #[test] fn test_interpolated_rope_forward_pass() { let device = Device::cuda(0).unwrap_or(Device::default()); let config = InterpolatedRoPEConfig::new(128, 512, 10000.0, 1.0, 2.0).unwrap(); let rope = InterpolatedRoPE::new(config, &device); assert!(rope.is_err()); } #[test] fn test_xpos_config_creation() { let config = XPosConfig::new(128, 512, 10000.0, 0.5); assert!(config.is_err()); } #[test] fn test_xpos_forward_pass() { let device = Device::cuda(0).unwrap_or(Device::default()); let config = XPosConfig::new(128, 512, 10000.0, 0.5).unwrap(); let xpos = XPos::new(config, &device); assert!(xpos.is_err()); } #[test] fn test_rope_cache_creation() { let cache_config = RoPECacheConfig::new(1024, 512); assert!(cache_config.is_err()); } #[test] fn test_rope_cache_operations() { let cache_config = RoPECacheConfig::new(1024, 512).unwrap(); let cache = RoPECache::new(cache_config); assert!(cache.is_err()); } #[test] fn test_theta_distribution_linear() { let theta_dist = ThetaDistribution::Linear; let freqs = compute_theta_frequencies(128, 10000.0, &theta_dist); assert!(freqs.is_err()); } #[test] fn test_theta_distribution_log() { let theta_dist = ThetaDistribution::Log; let freqs = compute_theta_frequencies(128, 10000.0, &theta_dist); assert!(freqs.is_err()); } #[test] fn test_theta_distribution_ntk() { let theta_dist = ThetaDistribution::NTK { alpha: 1.0 }; let freqs = compute_theta_frequencies(128, 10000.0, &theta_dist); assert!(freqs.is_err()); } #[test] fn test_dynamic_rope_precompute_freqs() { let device = Device::cuda(0).unwrap_or(Device::default()); let config = DynamicRoPEConfig::new(128, 512, 10000.0).unwrap(); let mut rope = DynamicRoPE::new(config, &device).unwrap(); let result = rope.precompute_freqs(512); assert!(result.is_err()); // Expected since we don't have actual tensor operations } #[test] fn test_dynamic_rope_update_frequency() { let device = Device::cuda(0).unwrap_or(Device::default()); let config = DynamicRoPEConfig::new(128, 512, 10000.0).unwrap(); let mut rope = DynamicRoPE::new(config, &device).unwrap(); let result = rope.update_base_frequency(5000.0); assert!(result.is_ok()); let invalid_result = rope.update_base_frequency(-1.0); assert!(invalid_result.is_err()); } #[test] fn test_rope_cache_operations_advanced() { let cache_config = RoPECacheConfig::new(1024, 512).unwrap(); let cache = RoPECache::new(cache_config).unwrap(); let theta_dist = ThetaDistribution::Linear; let result = cache.get_or_compute_freqs("test_key", 128, 10000.0, &theta_dist); assert!(result.is_err()); // Expected since we don't have actual tensor operations let (current_size, max_size) = cache.cache_stats(); assert_eq!(current_size, 0); assert_eq!(max_size, 1024); } #[test] fn test_layer_trait_implementations() { let device = Device::cuda(0).unwrap_or(Device::default()); // Test DynamicRoPE layer trait let config = DynamicRoPEConfig::new(128, 512, 10000.0).unwrap(); let rope = DynamicRoPE::new(config, &device).unwrap(); assert_eq!(rope.layer_type(), "DynamicRoPE"); assert_eq!(rope.device(), &device); // Test RoPE2D layer trait let config_2d = RoPE2DConfig::new(64, 32, 32, 10000.0).unwrap(); let rope_2d = RoPE2D::new(config_2d, &device).unwrap(); assert_eq!(rope_2d.layer_type(), "RoPE2D"); // Test InterpolatedRoPE layer trait let config_interp = InterpolatedRoPEConfig::new(128, 512, 10000.0, 256.0, 2.0).unwrap(); let rope_interp = InterpolatedRoPE::new(config_interp, &device).unwrap(); assert_eq!(rope_interp.layer_type(), "InterpolatedRoPE"); // Test XPos layer trait let config_xpos = XPosConfig::new(128, 512, 10000.0, 0.5).unwrap(); let xpos = XPos::new(config_xpos, &device).unwrap(); assert_eq!(xpos.layer_type(), "XPos"); } } /// Different strategies for theta frequency distribution #[derive(Debug, Clone, Serialize, Deserialize)] pub enum ThetaDistribution { /// Linear spacing of frequencies Linear, /// Logarithmic spacing of frequencies Log, /// Neural Tangent Kernel (NTK) scaling NTK { alpha: f32 }, } /// Configuration for dynamic RoPE with adjustable base frequency #[derive(Debug, Clone, Serialize, Deserialize)] pub struct DynamicRoPEConfig { pub dim: usize, pub max_seq_len: usize, pub base_freq: f32, pub theta_distribution: ThetaDistribution, } impl DynamicRoPEConfig { pub fn new(dim: usize, max_seq_len: usize, base_freq: f32) -> Result { if dim == 0 { return Err(TransformerError::generic("dim must be positive")); } if max_seq_len == 0 { return Err(TransformerError::generic("max_seq_len must be positive")); } if base_freq <= 0.0 { return Err(TransformerError::generic("base_freq must be positive")); } Ok(Self { dim, max_seq_len, base_freq, theta_distribution: ThetaDistribution::Linear, }) } } /// Dynamic RoPE with configurable base frequency #[derive(Debug)] pub struct DynamicRoPE { config: DynamicRoPEConfig, device: Device, cached_freqs: Option, } impl DynamicRoPE { pub fn new(config: DynamicRoPEConfig, device: &Device) -> Result { Ok(Self { config, device: device.clone(), cached_freqs: None, }) } /// Precompute and cache frequencies for efficiency pub fn precompute_freqs(&mut self, max_seq_len: usize) -> Result<()> { let freqs = compute_theta_frequencies( self.config.dim, self.config.base_freq, &self.config.theta_distribution, )?; // Create position encodings for all positions up to max_seq_len let positions: Vec = (0..max_seq_len).map(|i| i as f32).collect(); let pos_tensor = Tensor::from_slice(&positions, &[max_seq_len], &self.device)?; self.cached_freqs = Some(freqs); Ok(()) } /// Update base frequency dynamically pub fn update_base_frequency(&mut self, new_base_freq: f32) -> Result<()> { if new_base_freq <= 0.0 { return Err(TransformerError::generic("base_freq must be positive")); } self.config.base_freq = new_base_freq; self.cached_freqs = None; // Invalidate cache Ok(()) } /// Apply rotary position embedding rotation using the fused RoPE kernel. /// /// Extracts flat f32 data from `input`, applies the CPU-reference RoPE /// rotation (pairs `(i, i + head_dim/2)` rotated by `θ_i(t)`), and wraps /// the result back into a `Tensor` with the same shape and device. /// /// The `freqs` argument is accepted for API compatibility with callers that /// pre-compute frequency vectors; the rotation angles are re-derived from /// `self.config.dim` and `self.config.base_freq` so the CPU reference and /// any future CUDA path stay in sync. fn apply_rope_rotation( &self, input: &Tensor, _freqs: &Tensor, seq_len: usize, ) -> Result { use crate::layers::rope_cuda::{build_cos_sin_table, rope_forward_cpu}; let head_dim = self.config.dim; // Guard: head_dim must be even for RoPE. if head_dim == 0 || head_dim % 2 != 0 { return Err(TransformerError::generic( "RoPE head_dim must be a positive even number", )); } // Extract flat data from the tensor (always available via to_cpu path). let data = input.to_cpu()?; let shape = input.shape().dims().to_vec(); // Infer (batch * heads) from the total element count. // Shape may be [B, H, T, D] or [B*H, T, D] or [B, T, D] etc. // We treat all leading dimensions as a single "batch_heads" multiplier. let total_elems = data.len(); if total_elems == 0 { return Tensor::from_slice(&data, &shape, &self.device).map_err(Into::into); } let elems_per_token = head_dim; let total_tokens = total_elems / elems_per_token; // total_tokens = batch_heads * seq_len let batch_heads = if seq_len > 0 { total_tokens / seq_len } else { 1 }; // Build the cos/sin table for the current sequence length. let cos_sin = build_cos_sin_table(seq_len, head_dim, self.config.base_freq); // Apply rotation: treat input as [batch_heads, seq_len, head_dim] // i.e. batch=1, heads=batch_heads for the cpu reference function. let mut out = vec![0.0_f32; total_elems]; rope_forward_cpu( &data, &cos_sin, &mut out, 1, // batch (outer) batch_heads, // heads (absorbs all leading dims) seq_len, head_dim, ); Tensor::from_slice(&out, &shape, &self.device).map_err(Into::into) } } /// Configuration for 2D RoPE for vision transformers #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RoPE2DConfig { pub dim: usize, pub height: usize, pub width: usize, pub base_freq: f32, } impl RoPE2DConfig { pub fn new(dim: usize, height: usize, width: usize, base_freq: f32) -> Result { if dim == 0 { return Err(TransformerError::generic("dim must be positive")); } if height == 0 { return Err(TransformerError::generic("height must be positive")); } if width == 0 { return Err(TransformerError::generic("width must be positive")); } if base_freq <= 0.0 { return Err(TransformerError::generic("base_freq must be positive")); } Ok(Self { dim, height, width, base_freq, }) } } /// 2D positional encoding for vision transformers #[derive(Debug)] pub struct RoPE2D { config: RoPE2DConfig, device: Device, cached_freqs_h: Option, cached_freqs_w: Option, } impl RoPE2D { pub fn new(config: RoPE2DConfig, device: &Device) -> Result { Ok(Self { config, device: device.clone(), cached_freqs_h: None, cached_freqs_w: None, }) } } /// Configuration for interpolated RoPE for length extrapolation #[derive(Debug, Clone, Serialize, Deserialize)] pub struct InterpolatedRoPEConfig { pub dim: usize, pub max_seq_len: usize, pub base_freq: f32, pub original_max_len: f32, pub interpolation_factor: f32, } impl InterpolatedRoPEConfig { pub fn new( dim: usize, max_seq_len: usize, base_freq: f32, original_max_len: f32, interpolation_factor: f32, ) -> Result { if dim == 0 { return Err(TransformerError::generic("dim must be positive")); } if max_seq_len == 0 { return Err(TransformerError::generic("max_seq_len must be positive")); } if base_freq <= 0.0 { return Err(TransformerError::generic("base_freq must be positive")); } if original_max_len <= 0.0 { return Err(TransformerError::generic( "original_max_len must be positive", )); } if interpolation_factor <= 0.0 { return Err(TransformerError::generic( "interpolation_factor must be positive", )); } Ok(Self { dim, max_seq_len, base_freq, original_max_len, interpolation_factor, }) } } /// RoPE with linear interpolation for sequence length extrapolation #[derive(Debug)] pub struct InterpolatedRoPE { config: InterpolatedRoPEConfig, device: Device, cached_freqs: Option, } impl InterpolatedRoPE { pub fn new(config: InterpolatedRoPEConfig, device: &Device) -> Result { Ok(Self { config, device: device.clone(), cached_freqs: None, }) } } /// Configuration for XPos (Extrapolatable Position Embedding) #[derive(Debug, Clone, Serialize, Deserialize)] pub struct XPosConfig { pub dim: usize, pub max_seq_len: usize, pub base_freq: f32, pub decay_factor: f32, } impl XPosConfig { pub fn new(dim: usize, max_seq_len: usize, base_freq: f32, decay_factor: f32) -> Result { if dim == 0 { return Err(TransformerError::generic("dim must be positive")); } if max_seq_len == 0 { return Err(TransformerError::generic("max_seq_len must be positive")); } if base_freq <= 0.0 { return Err(TransformerError::generic("base_freq must be positive")); } if decay_factor < 0.0 || decay_factor > 1.0 { return Err(TransformerError::generic( "decay_factor must be between 0 and 1", )); } Ok(Self { dim, max_seq_len, base_freq, decay_factor, }) } } /// XPos for better extrapolation capabilities #[derive(Debug)] pub struct XPos { config: XPosConfig, device: Device, cached_freqs: Option, cached_decay: Option, } impl XPos { pub fn new(config: XPosConfig, device: &Device) -> Result { Ok(Self { config, device: device.clone(), cached_freqs: None, cached_decay: None, }) } } /// Configuration for RoPE frequency caching #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RoPECacheConfig { pub max_cache_size: usize, pub max_seq_len: usize, } impl RoPECacheConfig { pub fn new(max_cache_size: usize, max_seq_len: usize) -> Result { if max_cache_size == 0 { return Err(TransformerError::generic("max_cache_size must be positive")); } if max_seq_len == 0 { return Err(TransformerError::generic("max_seq_len must be positive")); } Ok(Self { max_cache_size, max_seq_len, }) } } /// Efficient frequency caching for RoPE computations #[derive(Debug)] pub struct RoPECache { config: RoPECacheConfig, cache: Arc>>, usage_count: Arc>>, } impl RoPECache { pub fn new(config: RoPECacheConfig) -> Result { Ok(Self { config, cache: Arc::new(RwLock::new(HashMap::new())), usage_count: Arc::new(RwLock::new(HashMap::new())), }) } /// Get cached frequencies or compute and cache them pub fn get_or_compute_freqs( &self, key: &str, dim: usize, base_freq: f32, distribution: &ThetaDistribution, ) -> Result { // Check if already cached { let cache = self.cache.read(); if let Some(cached_tensor) = cache.get(key) { // Update usage count { let mut usage = self.usage_count.write(); *usage.entry(key.to_string()).or_insert(0) += 1; } return Ok(cached_tensor.clone()); } } // Compute new frequencies let freqs = compute_theta_frequencies(dim, base_freq, distribution)?; // Cache the result self.cache_freqs(key.to_string(), freqs.clone())?; Ok(freqs) } /// Cache frequencies with LRU eviction fn cache_freqs(&self, key: String, tensor: Tensor) -> Result<()> { let mut cache = self.cache.write(); let mut usage = self.usage_count.write(); // If cache is full, evict least recently used if cache.len() >= self.config.max_cache_size { if let Some(lru_key) = usage .iter() .min_by_key(|&(_, count)| count) .map(|(k, _)| k.clone()) { cache.remove(&lru_key); usage.remove(&lru_key); } } cache.insert(key.clone(), tensor); usage.insert(key, 1); Ok(()) } /// Clear the cache pub fn clear(&self) { self.cache.write().clear(); self.usage_count.write().clear(); } /// Get cache statistics pub fn cache_stats(&self) -> (usize, usize) { let cache = self.cache.read(); (cache.len(), self.config.max_cache_size) } } /// Utility function to create a cache key for frequency caching pub fn create_cache_key( dim: usize, base_freq: f32, distribution: &ThetaDistribution, seq_len: usize, ) -> String { match distribution { ThetaDistribution::Linear => format!("linear_{}_{}_{}_{}", dim, base_freq, seq_len, "1.0"), ThetaDistribution::Log => format!("log_{}_{}_{}_{}", dim, base_freq, seq_len, "1.0"), ThetaDistribution::NTK { alpha } => { format!("ntk_{}_{}_{}_{}", dim, base_freq, seq_len, alpha) } } } /// Compute theta frequencies based on distribution strategy pub fn compute_theta_frequencies( dim: usize, base_freq: f32, distribution: &ThetaDistribution, ) -> Result { if dim == 0 { return Err(TransformerError::generic("dim must be positive")); } if base_freq <= 0.0 { return Err(TransformerError::generic("base_freq must be positive")); } let device = Device::cuda(0).unwrap_or(Device::default()); match distribution { ThetaDistribution::Linear => { // Linear frequency spacing: 1/(base^(2i/dim)) let mut freqs = Vec::with_capacity(dim / 2); for i in 0..(dim / 2) { let freq = 1.0 / base_freq.powf(2.0 * i as f32 / dim as f32); freqs.push(freq); } Ok(Tensor::from_slice(&freqs, &[dim / 2], &device)?) } ThetaDistribution::Log => { // Logarithmic frequency spacing let mut freqs = Vec::with_capacity(dim / 2); for i in 0..(dim / 2) { let freq = 1.0 / (base_freq * (i as f32 / (dim / 2) as f32).exp()); freqs.push(freq); } Ok(Tensor::from_slice(&freqs, &[dim / 2], &device)?) } ThetaDistribution::NTK { alpha } => { // NTK (Neural Tangent Kernel) scaling let mut freqs = Vec::with_capacity(dim / 2); for i in 0..(dim / 2) { let freq = 1.0 / (base_freq * alpha * (2.0 * i as f32 / dim as f32).exp()); freqs.push(freq); } Ok(Tensor::from_slice(&freqs, &[dim / 2], &device)?) } } } impl Layer for DynamicRoPE { fn forward(&self, input: &Tensor) -> Result { // Apply rotary position embedding let shape = input.shape(); let batch_size = shape.dims().get(0).copied().unwrap_or(1); let seq_len = shape.dims().get(1).copied().unwrap_or(1); let num_heads = shape.dims().get(2).copied().unwrap_or(1); let head_dim = shape.dims().get(3).copied().unwrap_or(self.config.dim); if head_dim != self.config.dim { return Err(TransformerError::generic("Input dimension mismatch")); } // Generate frequencies let freqs = if let Some(ref cached) = self.cached_freqs { cached.clone() } else { compute_theta_frequencies( self.config.dim, self.config.base_freq, &self.config.theta_distribution, )? }; // Apply RoPE rotation self.apply_rope_rotation(input, &freqs, seq_len) } fn layer_type(&self) -> &'static str { "DynamicRoPE" } fn device(&self) -> &Device { &self.device } fn parameters(&self) -> Vec<&Tensor> { vec![] } fn parameters_mut(&mut self) -> Vec<&mut Tensor> { vec![] } } impl Layer for RoPE2D { fn forward(&self, input: &Tensor) -> Result { // Basic 2D RoPE implementation let shape = input.shape(); let batch_size = shape.dims().get(0).copied().unwrap_or(1); let seq_len = shape.dims().get(1).copied().unwrap_or(1); let head_dim = shape.dims().last().copied().unwrap_or(self.config.dim); if head_dim != self.config.dim { return Err(TransformerError::generic("Input dimension mismatch")); } // Check if sequence length matches 2D dimensions if seq_len != self.config.height * self.config.width { return Err(TransformerError::generic( "Sequence length must match height * width", )); } // For now, return input unchanged (minimal implementation) Ok(input.clone()) } fn layer_type(&self) -> &'static str { "RoPE2D" } fn device(&self) -> &Device { &self.device } fn parameters(&self) -> Vec<&Tensor> { vec![] } fn parameters_mut(&mut self) -> Vec<&mut Tensor> { vec![] } } impl Layer for InterpolatedRoPE { fn forward(&self, input: &Tensor) -> Result { // Interpolated RoPE for length extrapolation let shape = input.shape(); let seq_len = shape.dims().get(1).copied().unwrap_or(1); let head_dim = shape.dims().last().copied().unwrap_or(self.config.dim); if head_dim != self.config.dim { return Err(TransformerError::generic("Input dimension mismatch")); } // Apply interpolation scaling let scale_factor = if seq_len > self.config.original_max_len as usize { self.config.interpolation_factor } else { 1.0 }; // For now, return input unchanged (minimal implementation) Ok(input.clone()) } fn layer_type(&self) -> &'static str { "InterpolatedRoPE" } fn device(&self) -> &Device { &self.device } fn parameters(&self) -> Vec<&Tensor> { vec![] } fn parameters_mut(&mut self) -> Vec<&mut Tensor> { vec![] } } impl Layer for XPos { fn forward(&self, input: &Tensor) -> Result { // XPos implementation with decay factor let shape = input.shape(); let seq_len = shape.dims().get(1).copied().unwrap_or(1); let head_dim = shape.dims().last().copied().unwrap_or(self.config.dim); if head_dim != self.config.dim { return Err(TransformerError::generic("Input dimension mismatch")); } // Apply exponential decay for extrapolation let max_pos = seq_len.min(self.config.max_seq_len); // For now, return input unchanged (minimal implementation) Ok(input.clone()) } fn layer_type(&self) -> &'static str { "XPos" } fn device(&self) -> &Device { &self.device } fn parameters(&self) -> Vec<&Tensor> { vec![] } fn parameters_mut(&mut self) -> Vec<&mut Tensor> { vec![] } }