Initial commit
This commit is contained in:
@@ -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(¶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<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(¤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<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(¶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<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(())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user