Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,726 @@
//! 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<MergedModel> {
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<IndexMap<String, crate::types::ParameterTensor>> {
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(&param1.data, &param2.data, t)?
} else {
self.spherical_interpolation(&param1.data, &param2.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<IndexMap<String, crate::types::ParameterTensor>> {
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(&current_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<Vec<f32>> {
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<Vec<f32>> {
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<f32>,
q2: &Vector4<f32>,
t: f32,
) -> Result<Vector4<f32>> {
// 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<f32> {
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<QualityMetrics> {
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::<f32>() / 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(&param.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<f32> = 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<MergedModel> {
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<f32>)>) -> 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(())
}
}