727 lines
25 KiB
Rust
727 lines
25 KiB
Rust
//! 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(())
|
|
}
|
|
}
|