//! Advanced model merging strategies //! //! This module implements sophisticated merging strategies that build upon //! the core algorithms: //! - Model Soups: Ensemble averaging with greedy selection //! - Frankenmerging: Layer-wise merging across architectures //! - Progressive: Gradual merging with validation checkpoints pub mod frankenmerge; pub mod model_soup; pub mod progressive; // Re-export for convenience pub use frankenmerge::*; pub use model_soup::*; pub use progressive::*; use crate::error::Result; use crate::types::{MergedModel, Model}; use std::collections::HashMap; /// Common utilities for advanced merging strategies pub struct StrategyUtils; impl StrategyUtils { /// Validate that all models in a group have compatible architectures pub fn validate_model_group(models: &[Model]) -> Result<()> { if models.len() < 2 { return Ok(()); } let reference = &models[0]; for (i, model) in models.iter().enumerate().skip(1) { if !reference.is_compatible_with(model) { return Err(crate::error::MergeError::compatibility(format!( "Model {i} is incompatible with reference model in group" ))); } } Ok(()) } /// Compute pairwise similarities between models pub fn compute_pairwise_similarities(models: &[Model]) -> Result>> { let n = models.len(); let mut similarities = vec![vec![0.0; n]; n]; for i in 0..n { similarities[i][i] = 1.0; // Self-similarity for j in (i + 1)..n { let similarity = crate::algorithms::MergeUtils::compute_similarity(&models[i], &models[j])?; similarities[i][j] = similarity; similarities[j][i] = similarity; // Symmetric } } Ok(similarities) } /// Group models by similarity threshold pub fn group_models_by_similarity(models: &[Model], threshold: f32) -> Result>> { let similarities = Self::compute_pairwise_similarities(models)?; let n = models.len(); let mut groups = Vec::new(); let mut assigned = vec![false; n]; for i in 0..n { if assigned[i] { continue; } let mut group = vec![i]; assigned[i] = true; // Find all models similar to model i for j in (i + 1)..n { if !assigned[j] && similarities[i][j] >= threshold { group.push(j); assigned[j] = true; } } groups.push(group); } Ok(groups) } /// Compute model diversity within a group pub fn compute_group_diversity(models: &[Model], group_indices: &[usize]) -> Result { if group_indices.len() < 2 { return Ok(0.0); } let group_models: Vec<&Model> = group_indices.iter().map(|&i| &models[i]).collect(); let similarities = Self::compute_pairwise_similarities( &group_models.into_iter().cloned().collect::>(), )?; let mut total_dissimilarity = 0.0; let mut pair_count = 0; for i in 0..similarities.len() { for j in (i + 1)..similarities.len() { total_dissimilarity += 1.0 - similarities[i][j]; pair_count += 1; } } Ok(if pair_count > 0 { total_dissimilarity / pair_count as f32 } else { 0.0 }) } /// Select representative models from a group pub fn select_representatives( models: &[Model], group_indices: &[usize], max_representatives: usize, ) -> Result> { if group_indices.len() <= max_representatives { return Ok(group_indices.to_vec()); } // Use k-means-like selection based on parameter space let mut representatives = Vec::new(); let mut remaining: Vec = group_indices.to_vec(); // First representative: model with highest average similarity to others let mut best_idx = 0; let mut best_avg_sim = 0.0; for (i, &model_idx) in remaining.iter().enumerate() { let mut total_sim = 0.0; for &other_idx in &remaining { if model_idx != other_idx { total_sim += crate::algorithms::MergeUtils::compute_similarity( &models[model_idx], &models[other_idx], )?; } } let avg_sim = total_sim / (remaining.len() - 1) as f32; if avg_sim > best_avg_sim { best_avg_sim = avg_sim; best_idx = i; } } representatives.push(remaining.remove(best_idx)); // Subsequent representatives: maximize diversity while representatives.len() < max_representatives && !remaining.is_empty() { let mut best_idx = 0; let mut best_min_sim = 1.0f32; for (i, &candidate_idx) in remaining.iter().enumerate() { let mut min_sim = 1.0f32; for &rep_idx in &representatives { let sim = crate::algorithms::MergeUtils::compute_similarity( &models[candidate_idx], &models[rep_idx], )?; min_sim = min_sim.min(sim); } if min_sim < best_min_sim { best_min_sim = min_sim; best_idx = i; } } representatives.push(remaining.remove(best_idx)); } Ok(representatives) } /// Compute model performance scores (placeholder - would use actual validation) pub fn compute_performance_scores(models: &[Model]) -> HashMap { let mut scores = HashMap::new(); for model in models { // Placeholder scoring based on model complexity let complexity = model.parameter_count() as f32; let diversity = Self::compute_parameter_diversity(model); // Simple heuristic score (in practice would use validation metrics) let score = (1.0 / (1.0 + (complexity / 1000000.0))).max(0.1) * (0.5 + 0.5 * diversity); scores.insert(model.name.clone(), score); } scores } /// Compute parameter diversity within a single model pub fn compute_parameter_diversity(model: &Model) -> f32 { let all_params: Vec = model .parameters .values() .flat_map(|p| p.data.iter()) .copied() .collect(); if all_params.is_empty() { return 0.0; } let (_, std_dev, _, _) = crate::algorithms::MergeUtils::compute_moments(&all_params); std_dev.min(1.0) } /// Create validation checkpoints for progressive merging pub fn create_validation_checkpoints( base_model: &Model, target_model: &Model, num_checkpoints: usize, ) -> Result> { let mut checkpoints = Vec::new(); for i in 0..=num_checkpoints { let t = i as f32 / num_checkpoints as f32; // Linear interpolation between base and target let mut checkpoint_model = base_model.clone(); checkpoint_model.id = uuid::Uuid::new_v4(); checkpoint_model.name = format!("checkpoint_{i}_{t:.2}"); // Interpolate parameters for (param_name, base_param) in &base_model.parameters { if let Some(target_param) = target_model.parameters.get(param_name) { let interpolated_data = crate::algorithms::MergeUtils::elementwise_op( &base_param.data, &target_param.data, |base_val, target_val| (1.0 - t) * base_val + t * target_val, )?; if let Some(checkpoint_param) = checkpoint_model.parameters.get_mut(param_name) { checkpoint_param.data = interpolated_data; } } } // Create merged model wrapper let merged_checkpoint = MergedModel { model: checkpoint_model, merge_info: crate::types::MergeInfo { strategy: "ValidationCheckpoint".to_string(), source_models: vec![ crate::types::ModelReference { id: base_model.id, name: base_model.name.clone(), path: std::path::PathBuf::from(&base_model.name), weight: Some(1.0 - t), }, crate::types::ModelReference { id: target_model.id, name: target_model.name.clone(), path: std::path::PathBuf::from(&target_model.name), weight: Some(t), }, ], merged_at: chrono::Utc::now(), config: serde_json::json!({"t": t, "checkpoint": i}), statistics: crate::types::MergeStatistics::default(), }, quality_metrics: crate::types::QualityMetrics::default(), }; checkpoints.push(merged_checkpoint); } Ok(checkpoints) } /// Estimate memory requirements for merging strategy pub fn estimate_memory_requirements(models: &[Model], strategy: &str) -> Result { let total_model_memory: usize = models.iter().map(super::types::Model::memory_size).sum(); let overhead_factor = match strategy { "ModelSoup" => 1.2, // Small overhead for ensemble averaging "Frankenmerge" => 1.5, // Moderate overhead for layer-wise operations "Progressive" => 2.0, // Higher overhead for checkpoints _ => 1.1, }; let estimated_memory = (total_model_memory as f32 * overhead_factor) as usize; Ok(estimated_memory) } /// Check resource constraints pub fn check_resource_constraints( models: &[Model], strategy: &str, max_memory_mb: Option, ) -> Result<()> { let estimated_memory = Self::estimate_memory_requirements(models, strategy)?; let estimated_memory_mb = estimated_memory / (1024 * 1024); if let Some(max_mb) = max_memory_mb && estimated_memory_mb > max_mb { return Err(crate::error::MergeError::memory(estimated_memory_mb)); } // Check for extremely large memory requirements if estimated_memory_mb > 32768 { // 32GB return Err(crate::error::MergeError::resource_exhausted(format!( "Memory requirement: {estimated_memory_mb}MB" ))); } Ok(()) } } #[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: std::collections::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 } #[test] fn test_model_group_validation() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0])]); let model2 = create_test_model("model2", vec![("weight".to_string(), vec![3.0, 4.0])]); let result = StrategyUtils::validate_model_group(&[model1, model2]); assert!(result.is_ok()); Ok(()) } #[test] fn test_pairwise_similarities() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0, 3.0])]); let model2 = create_test_model( "model2", vec![ ("weight".to_string(), vec![1.0, 2.0, 3.0]), // Identical ], ); let model3 = create_test_model( "model3", vec![ ("weight".to_string(), vec![4.0, 5.0, 6.0]), // Different ], ); let similarities = StrategyUtils::compute_pairwise_similarities(&[model1, model2, model3])?; assert_eq!(similarities.len(), 3); assert_eq!(similarities[0].len(), 3); // Self-similarities should be 1.0 assert_eq!(similarities[0][0], 1.0); assert_eq!(similarities[1][1], 1.0); assert_eq!(similarities[2][2], 1.0); // Model1 and Model2 should be very similar (identical) assert!(similarities[0][1] > 0.9); // Model1/Model2 and Model3 should be less similar assert!(similarities[0][2] < similarities[0][1]); Ok(()) } #[test] fn test_similarity_grouping() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0])]); let model2 = create_test_model( "model2", vec![ ("weight".to_string(), vec![1.1, 2.1]), // Very similar to model1 ], ); let model3 = create_test_model( "model3", vec![ ("weight".to_string(), vec![10.0, 20.0]), // Very different ], ); let groups = StrategyUtils::group_models_by_similarity(&[model1, model2, model3], 0.8)?; // Should have at least 2 groups (similar models together, different model separate) assert!(groups.len() >= 2); // Total models should be preserved let total_models: usize = groups.iter().map(|g| g.len()).sum(); assert_eq!(total_models, 3); Ok(()) } #[test] fn test_representative_selection() -> Result<()> { let models = (0..5) .map(|i| { create_test_model( &format!("model{}", i), vec![("weight".to_string(), vec![i as f32, (i + 1) as f32])], ) }) .collect::>(); let group_indices = vec![0, 1, 2, 3, 4]; let representatives = StrategyUtils::select_representatives(&models, &group_indices, 3)?; assert_eq!(representatives.len(), 3); assert!(representatives.iter().all(|&i| i < 5)); // Should select diverse representatives assert!(representatives[0] != representatives[1]); assert!(representatives[1] != representatives[2]); assert!(representatives[0] != representatives[2]); Ok(()) } #[test] fn test_parameter_diversity() { let model1 = create_test_model( "uniform", vec![ ("weight".to_string(), vec![1.0, 1.0, 1.0, 1.0]), // No diversity ], ); let model2 = create_test_model( "diverse", vec![ ("weight".to_string(), vec![1.0, 5.0, 10.0, 2.0]), // High diversity ], ); let diversity1 = StrategyUtils::compute_parameter_diversity(&model1); let diversity2 = StrategyUtils::compute_parameter_diversity(&model2); assert!(diversity1 < diversity2); assert_eq!(diversity1, 0.0); // Uniform values should have zero diversity assert!(diversity2 > 0.0); // Diverse values should have positive diversity } #[test] fn test_memory_estimation() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0; 1000])]); let model2 = create_test_model("model2", vec![("weight".to_string(), vec![1.0; 1000])]); let models = vec![model1, model2]; let soup_memory = StrategyUtils::estimate_memory_requirements(&models, "ModelSoup")?; let progressive_memory = StrategyUtils::estimate_memory_requirements(&models, "Progressive")?; // Progressive should require more memory than ModelSoup assert!(progressive_memory > soup_memory); Ok(()) } #[test] fn test_resource_constraints() -> Result<()> { let small_model = create_test_model("small", vec![("weight".to_string(), vec![1.0; 10])]); let models = vec![small_model]; // Should pass with reasonable memory limit let result = StrategyUtils::check_resource_constraints(&models, "ModelSoup", Some(1024)); assert!(result.is_ok()); // Should fail with very small memory limit let result = StrategyUtils::check_resource_constraints(&models, "ModelSoup", Some(1)); assert!(result.is_err()); Ok(()) } #[test] fn test_validation_checkpoints() -> Result<()> { let base_model = create_test_model("base", vec![("weight".to_string(), vec![0.0, 0.0])]); let target_model = create_test_model("target", vec![("weight".to_string(), vec![2.0, 4.0])]); let checkpoints = StrategyUtils::create_validation_checkpoints(&base_model, &target_model, 4)?; assert_eq!(checkpoints.len(), 5); // 0, 1, 2, 3, 4 (inclusive) // First checkpoint should be base model let first_weight = checkpoints[0].model.get_parameter("weight").unwrap(); assert_eq!(first_weight.data, vec![0.0, 0.0]); // Last checkpoint should be target model let last_weight = checkpoints[4].model.get_parameter("weight").unwrap(); assert_eq!(last_weight.data, vec![2.0, 4.0]); // Middle checkpoint should be interpolated let middle_weight = checkpoints[2].model.get_parameter("weight").unwrap(); assert_eq!(middle_weight.data, vec![1.0, 2.0]); // 50% interpolation Ok(()) } }