//! SLERP (Spherical Linear Interpolation) algorithm implementation //! //! SLERP merges models by: //! 1. Normalizing parameter vectors to unit spheres //! 2. Performing spherical interpolation between vectors //! 3. Optionally using quaternion-based interpolation for rotational parameters //! 4. Supporting adaptive interpolation weights //! //! Reference: "Spherical Linear Interpolation for Neural Network Merging" use crate::algorithms::MergeUtils; use crate::config::{NormalizationMethod, SlerpConfig}; use crate::error::{MergeError, Result}; use crate::types::{ MergeInfo, MergeStatistics, MergedModel, Model, ModelReference, QualityMetrics, }; use indexmap::IndexMap; use nalgebra::Vector4; use std::collections::HashMap; use tracing::{debug, info, warn}; /// SLERP merging algorithm implementation pub struct SlerpMerger { config: SlerpConfig, } impl SlerpMerger { /// Create a new SLERP merger with configuration pub fn new(config: SlerpConfig) -> Self { Self { config } } /// Execute SLERP merging on multiple models pub async fn merge(&self, models: &[Model]) -> Result { if models.len() < 2 { return Err(MergeError::algorithm( "SLERP", "At least 2 models required for merging", )); } if models.len() > 2 && !self.config.adaptive_t { warn!( "SLERP typically works with 2 models. Using pairwise interpolation for {} models", models.len() ); } info!( "Starting SLERP merge of {} models with t={}", models.len(), self.config.t ); let start_time = std::time::Instant::now(); // Step 1: Validate model compatibility self.validate_compatibility(models)?; // Step 2: Apply SLERP algorithm let merged_parameters = if models.len() == 2 { self.slerp_two_models(&models[0], &models[1]).await? } else { self.slerp_multiple_models(models).await? }; // Step 3: Create merged model let mut merged_model = models[0].clone(); merged_model.id = uuid::Uuid::new_v4(); merged_model.name = format!("SLERP_merged_{}", models.len()); merged_model.parameters = merged_parameters; // Step 4: Compute merge statistics and quality metrics let statistics = self.compute_statistics(&merged_model, models, start_time.elapsed()); let quality_metrics = self.compute_quality_metrics(&merged_model, models)?; let merge_info = MergeInfo { strategy: "SLERP".to_string(), source_models: models .iter() .enumerate() .map(|(i, m)| ModelReference { id: m.id, name: m.name.clone(), path: std::path::PathBuf::from(&m.name), weight: Some(if i == 0 { 1.0 - self.config.t } else { self.config.t }), }) .collect(), merged_at: chrono::Utc::now(), config: serde_json::to_value(&self.config) .map_err(|e| MergeError::internal(format!("Config serialization: {e}")))?, statistics, }; info!("SLERP merge completed in {:?}", start_time.elapsed()); Ok(MergedModel { model: merged_model, merge_info, quality_metrics, }) } /// Validate that models are compatible for SLERP merging fn validate_compatibility(&self, models: &[Model]) -> Result<()> { let base_model = &models[0]; for (i, model) in models.iter().enumerate().skip(1) { if !base_model.is_compatible_with(model) { return Err(MergeError::compatibility(format!( "Model {i} incompatible with base model" ))); } } Ok(()) } /// Perform SLERP interpolation between two models async fn slerp_two_models( &self, model1: &Model, model2: &Model, ) -> Result> { info!("Performing SLERP interpolation between two models"); let mut merged_parameters = IndexMap::new(); for (param_name, param1) in &model1.parameters { if let Some(param2) = model2.parameters.get(param_name) { debug!("Interpolating parameter: {}", param_name); let t = self.get_parameter_weight(param_name); let interpolated_data = if self.config.use_quaternions && self.is_rotational_parameter(param_name) { self.quaternion_slerp(¶m1.data, ¶m2.data, t)? } else { self.spherical_interpolation(¶m1.data, ¶m2.data, t)? }; let mut merged_param = param1.clone(); merged_param.data = interpolated_data; merged_parameters.insert(param_name.clone(), merged_param); } else { // Parameter only exists in model1, keep as is merged_parameters.insert(param_name.clone(), param1.clone()); } } // Add parameters that only exist in model2 for (param_name, param2) in &model2.parameters { if !model1.parameters.contains_key(param_name) { warn!( "Parameter {} only exists in model2, adding with weight t={}", param_name, self.config.t ); let mut scaled_param = param2.clone(); for value in &mut scaled_param.data { *value *= self.config.t; } merged_parameters.insert(param_name.clone(), scaled_param); } } Ok(merged_parameters) } /// Perform SLERP interpolation between multiple models using pairwise approach async fn slerp_multiple_models( &self, models: &[Model], ) -> Result> { info!( "Performing pairwise SLERP interpolation for {} models", models.len() ); // Start with the first model let mut current_model = models[0].clone(); // Progressively interpolate with each subsequent model for (i, next_model) in models.iter().enumerate().skip(1) { let weight = if self.config.adaptive_t { // Adaptive weight that gives equal influence to all models 1.0 / (i as f32 + 1.0) } else { self.config.t }; info!("Interpolating with model {} using weight: {}", i, weight); // Create temporary config with updated weight let temp_config = SlerpConfig { t: weight, ..self.config.clone() }; let temp_merger = Self::new(temp_config); let intermediate_params = temp_merger .slerp_two_models(¤t_model, next_model) .await?; // Update current model with interpolated parameters current_model.parameters = intermediate_params; } Ok(current_model.parameters) } /// Get interpolation weight for a specific parameter fn get_parameter_weight(&self, param_name: &str) -> f32 { if let Some(ref weights) = self.config.parameter_weights { weights.get(param_name).copied().unwrap_or(self.config.t) } else { self.config.t } } /// Check if a parameter represents rotational data (heuristic) fn is_rotational_parameter(&self, param_name: &str) -> bool { param_name.contains("rotation") || param_name.contains("quaternion") || param_name.contains("angle") || (param_name.contains("weight") && param_name.contains("attention")) } /// Perform spherical linear interpolation between two parameter vectors fn spherical_interpolation(&self, vec1: &[f32], vec2: &[f32], t: f32) -> Result> { if vec1.len() != vec2.len() { return Err(MergeError::parameter_mismatch( format!("{} elements", vec1.len()), format!("{} elements", vec2.len()), )); } if !(0.0..=1.0).contains(&t) { return Err(MergeError::numerical(format!( "Invalid interpolation parameter t: {t}" ))); } // Handle edge cases if t == 0.0 { return Ok(vec1.to_vec()); } if t == 1.0 { return Ok(vec2.to_vec()); } // Normalize vectors let mut norm_vec1 = vec1.to_vec(); let mut norm_vec2 = vec2.to_vec(); match self.config.normalization { NormalizationMethod::L2 => { MergeUtils::normalize_l2(&mut norm_vec1)?; MergeUtils::normalize_l2(&mut norm_vec2)?; } NormalizationMethod::L1 => { MergeUtils::normalize_l1(&mut norm_vec1)?; MergeUtils::normalize_l1(&mut norm_vec2)?; } NormalizationMethod::None => { // No normalization } } // Compute cosine of angle between vectors let dot_product = MergeUtils::cosine_similarity(&norm_vec1, &norm_vec2)?; let omega = dot_product.acos(); // Handle parallel vectors (avoid division by zero) if omega.sin().abs() < 1e-6 { warn!("Vectors are nearly parallel, falling back to linear interpolation"); return Ok(self.linear_interpolation(vec1, vec2, t)); } // Perform SLERP let sin_omega = omega.sin(); let factor1 = ((1.0 - t) * omega).sin() / sin_omega; let factor2 = (t * omega).sin() / sin_omega; let mut result = Vec::with_capacity(vec1.len()); for i in 0..vec1.len() { let interpolated = factor1 * norm_vec1[i] + factor2 * norm_vec2[i]; // Scale back using original vector magnitudes let original_scale = (1.0 - t) * vec1[i].abs() + t * vec2[i].abs(); result.push(interpolated * original_scale); } Ok(result) } /// Perform quaternion-based SLERP for rotational parameters fn quaternion_slerp(&self, vec1: &[f32], vec2: &[f32], t: f32) -> Result> { if vec1.len() != vec2.len() { return Err(MergeError::parameter_mismatch( format!("{} elements", vec1.len()), format!("{} elements", vec2.len()), )); } // Process parameters in groups of 4 (quaternions) or use regular SLERP if vec1.len().is_multiple_of(4) { let mut result = Vec::with_capacity(vec1.len()); for chunk_start in (0..vec1.len()).step_by(4) { let chunk_end = std::cmp::min(chunk_start + 4, vec1.len()); if chunk_end - chunk_start == 4 { // Full quaternion let q1 = Vector4::new( vec1[chunk_start], vec1[chunk_start + 1], vec1[chunk_start + 2], vec1[chunk_start + 3], ); let q2 = Vector4::new( vec2[chunk_start], vec2[chunk_start + 1], vec2[chunk_start + 2], vec2[chunk_start + 3], ); let interpolated = self.slerp_quaternions(&q1, &q2, t)?; result.extend_from_slice(interpolated.as_slice()); } else { // Partial quaternion, use regular spherical interpolation let chunk1 = &vec1[chunk_start..chunk_end]; let chunk2 = &vec2[chunk_start..chunk_end]; let interpolated = self.spherical_interpolation(chunk1, chunk2, t)?; result.extend(interpolated); } } Ok(result) } else { // Not divisible by 4, use regular spherical interpolation self.spherical_interpolation(vec1, vec2, t) } } /// SLERP for unit quaternions fn slerp_quaternions( &self, q1: &Vector4, q2: &Vector4, t: f32, ) -> Result> { // Normalize quaternions let norm1 = q1.norm(); let norm2 = q2.norm(); if norm1 == 0.0 || norm2 == 0.0 { return Err(MergeError::numerical("Zero-norm quaternion")); } let q1_unit = q1 / norm1; let q2_unit = q2 / norm2; // Compute dot product let mut dot = q1_unit.dot(&q2_unit); // If dot product is negative, negate one quaternion to take shorter path let q2_adjusted = if dot < 0.0 { dot = -dot; -q2_unit } else { q2_unit }; // If quaternions are very close, use linear interpolation if dot > 0.9995 { let result = q1_unit * (1.0 - t) + q2_adjusted * t; let result_norm = result.norm(); if result_norm == 0.0 { return Err(MergeError::numerical( "Zero-norm result in quaternion interpolation", )); } return Ok(result / result_norm * ((1.0 - t) * norm1 + t * norm2)); } // Calculate angle and perform spherical interpolation let theta_0 = dot.acos(); let sin_theta_0 = theta_0.sin(); let theta = theta_0 * t; let q2_perp = (q2_adjusted - q1_unit * dot) / sin_theta_0; let result = q1_unit * theta.cos() + q2_perp * theta.sin(); // Scale back to original magnitude let final_magnitude = (1.0 - t) * norm1 + t * norm2; Ok(result * final_magnitude) } /// Fallback linear interpolation fn linear_interpolation(&self, vec1: &[f32], vec2: &[f32], t: f32) -> Vec { vec1.iter() .zip(vec2.iter()) .map(|(&v1, &v2)| (1.0 - t) * v1 + t * v2) .collect() } /// Compute merge statistics fn compute_statistics( &self, merged_model: &Model, _source_models: &[Model], duration: std::time::Duration, ) -> MergeStatistics { let parameters_merged = merged_model.parameter_count(); let memory_usage_mb = merged_model.memory_size() / (1024 * 1024); MergeStatistics { parameters_merged, parameters_conflicted: 0, // SLERP doesn't have explicit conflicts parameters_dropped: 0, // SLERP doesn't drop parameters duration_ms: duration.as_millis() as u64, memory_usage_mb, gpu_memory_mb: None, } } /// Compute quality metrics for the merged model fn compute_quality_metrics( &self, merged_model: &Model, source_models: &[Model], ) -> Result { let mut consistency_scores = Vec::new(); // Compute consistency with each source model for source in source_models { let similarity = MergeUtils::compute_similarity(merged_model, source)?; consistency_scores.push(similarity); } let consistency_score = consistency_scores.iter().sum::() / consistency_scores.len() as f32; // Compute smoothness as a quality indicator let mut total_smoothness = 0.0; let mut param_count = 0; for param in merged_model.parameters.values() { if param.data.len() > 1 { // Compute variance as inverse smoothness measure let (_, std_dev, _, _) = MergeUtils::compute_moments(¶m.data); total_smoothness += 1.0 / (1.0 + std_dev); param_count += 1; } } let smoothness_score = if param_count > 0 { total_smoothness / param_count as f32 } else { 0.0 }; // Compute complexity score based on parameter distribution let all_params: Vec = merged_model .parameters .values() .flat_map(|p| p.data.iter()) .copied() .collect(); let (_, std_dev, _, _) = MergeUtils::compute_moments(&all_params); let complexity_score = std_dev.min(1.0); let mut quality_indicators = HashMap::new(); quality_indicators.insert("smoothness_score".to_string(), smoothness_score); quality_indicators.insert("parameter_diversity".to_string(), std_dev); quality_indicators.insert("average_similarity".to_string(), consistency_score); quality_indicators.insert( "interpolation_quality".to_string(), consistency_score * smoothness_score, ); let mut validation_results = HashMap::new(); validation_results.insert("smooth_interpolation".to_string(), smoothness_score > 0.5); validation_results.insert( "consistency_maintained".to_string(), consistency_score > 0.7, ); validation_results.insert("spherical_properties".to_string(), complexity_score > 0.05); let mut performance_predictions = HashMap::new(); performance_predictions.insert("expected_accuracy".to_string(), consistency_score * 0.95); performance_predictions.insert("stability_score".to_string(), smoothness_score * 0.9); performance_predictions.insert( "generalization".to_string(), (consistency_score + smoothness_score) * 0.5, ); Ok(QualityMetrics { consistency_score, complexity_score, quality_indicators, validation_results, performance_predictions, }) } } /// Convenience function for SLERP merging pub async fn merge_models(models: &[Model], config: &SlerpConfig) -> Result { let merger = SlerpMerger::new(config.clone()); merger.merge(models).await } #[cfg(test)] mod tests { use super::*; use crate::types::{DataType, ModelArchitecture, ParameterTensor}; fn create_test_model(name: &str, params: Vec<(String, Vec)>) -> Model { let arch = ModelArchitecture { arch_type: "test".to_string(), num_layers: 1, hidden_dim: params.len(), params: HashMap::new(), }; let mut model = Model::new(name.to_string(), arch); for (param_name, data) in params { let param = ParameterTensor::new(param_name, vec![data.len()], DataType::Float32, data); model.add_parameter(param); } model } #[tokio::test] async fn test_slerp_merge_basic() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 0.0, 0.0])]); let model2 = create_test_model("model2", vec![("weight".to_string(), vec![0.0, 1.0, 0.0])]); let config = SlerpConfig { t: 0.5, normalization: NormalizationMethod::L2, ..SlerpConfig::default() }; let merger = SlerpMerger::new(config); let result = merger.merge(&[model1, model2]).await?; assert!(result.model.parameters.contains_key("weight")); assert_eq!(result.merge_info.strategy, "SLERP"); let merged_weight = result.model.get_parameter("weight").unwrap(); assert_eq!(merged_weight.data.len(), 3); Ok(()) } #[test] fn test_spherical_interpolation() -> Result<()> { let merger = SlerpMerger::new(SlerpConfig::default()); let vec1 = vec![1.0, 0.0, 0.0]; let vec2 = vec![0.0, 1.0, 0.0]; let result = merger.spherical_interpolation(&vec1, &vec2, 0.5)?; assert_eq!(result.len(), 3); // Result should be somewhere between the two vectors assert!(result[0] > 0.0 && result[0] < 1.0); assert!(result[1] > 0.0 && result[1] < 1.0); Ok(()) } #[test] fn test_spherical_interpolation_edge_cases() -> Result<()> { let merger = SlerpMerger::new(SlerpConfig::default()); let vec1 = vec![1.0, 2.0, 3.0]; let vec2 = vec![4.0, 5.0, 6.0]; // t = 0 should return vec1 let result_0 = merger.spherical_interpolation(&vec1, &vec2, 0.0)?; assert_eq!(result_0, vec1); // t = 1 should return vec2 let result_1 = merger.spherical_interpolation(&vec1, &vec2, 1.0)?; assert_eq!(result_1, vec2); Ok(()) } #[test] fn test_quaternion_slerp() -> Result<()> { let merger = SlerpMerger::new(SlerpConfig { use_quaternions: true, ..SlerpConfig::default() }); // Two unit quaternions let q1 = Vector4::new(1.0, 0.0, 0.0, 0.0); let q2 = Vector4::new(0.0, 1.0, 0.0, 0.0); let result = merger.slerp_quaternions(&q1, &q2, 0.5)?; // Result should be normalized let norm = result.norm(); assert!((norm - 1.0).abs() < 1e-5); Ok(()) } #[tokio::test] async fn test_slerp_multiple_models() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 0.0])]); let model2 = create_test_model("model2", vec![("weight".to_string(), vec![0.0, 1.0])]); let model3 = create_test_model("model3", vec![("weight".to_string(), vec![-1.0, 0.0])]); let config = SlerpConfig { t: 0.5, adaptive_t: true, ..SlerpConfig::default() }; let merger = SlerpMerger::new(config); let result = merger.merge(&[model1, model2, model3]).await?; assert!(result.model.parameters.contains_key("weight")); let merged_weight = result.model.get_parameter("weight").unwrap(); assert_eq!(merged_weight.data.len(), 2); Ok(()) } #[test] fn test_parameter_weight_lookup() { let mut param_weights = HashMap::new(); param_weights.insert("attention.weight".to_string(), 0.3); param_weights.insert("mlp.weight".to_string(), 0.7); let config = SlerpConfig { t: 0.5, parameter_weights: Some(param_weights), ..SlerpConfig::default() }; let merger = SlerpMerger::new(config); assert_eq!(merger.get_parameter_weight("attention.weight"), 0.3); assert_eq!(merger.get_parameter_weight("mlp.weight"), 0.7); assert_eq!(merger.get_parameter_weight("unknown.weight"), 0.5); // Default t } #[test] fn test_rotational_parameter_detection() { let merger = SlerpMerger::new(SlerpConfig::default()); assert!(merger.is_rotational_parameter("attention.rotation_weight")); assert!(merger.is_rotational_parameter("layer.quaternion_param")); assert!(merger.is_rotational_parameter("angle_embedding")); assert!(merger.is_rotational_parameter("attention.weight")); assert!(!merger.is_rotational_parameter("linear.bias")); assert!(!merger.is_rotational_parameter("norm.scale")); } #[test] fn test_linear_interpolation_fallback() { let merger = SlerpMerger::new(SlerpConfig::default()); let vec1 = vec![1.0, 2.0, 3.0]; let vec2 = vec![4.0, 5.0, 6.0]; let result = merger.linear_interpolation(&vec1, &vec2, 0.5); let expected = vec![2.5, 3.5, 4.5]; assert_eq!(result, expected); } #[test] fn test_normalization_methods() -> Result<()> { let vec1 = vec![3.0, 4.0]; // |vec| = 5 // L2 normalization let config_l2 = SlerpConfig { normalization: NormalizationMethod::L2, ..SlerpConfig::default() }; let merger_l2 = SlerpMerger::new(config_l2); let vec2 = vec![0.0, 5.0]; // |vec| = 5 let result_l2 = merger_l2.spherical_interpolation(&vec1, &vec2, 0.5)?; assert!(result_l2.len() == 2); // L1 normalization let config_l1 = SlerpConfig { normalization: NormalizationMethod::L1, ..SlerpConfig::default() }; let merger_l1 = SlerpMerger::new(config_l1); let result_l1 = merger_l1.spherical_interpolation(&vec1, &vec2, 0.5)?; assert!(result_l1.len() == 2); // No normalization let config_none = SlerpConfig { normalization: NormalizationMethod::None, ..SlerpConfig::default() }; let merger_none = SlerpMerger::new(config_none); let result_none = merger_none.spherical_interpolation(&vec1, &vec2, 0.5)?; assert!(result_none.len() == 2); Ok(()) } }