//! Model Soup strategy implementation //! //! Model Soup creates ensemble models by: //! 1. Averaging parameters across multiple models //! 2. Optionally using greedy soup construction for optimal subset selection //! 3. Supporting weighted averaging based on model performance //! 4. Validating improvements at each step //! //! Reference: "Model soups: averaging weights of multiple fine-tuned models improves accuracy without increasing inference time" (Wortsman et al., 2022) use crate::algorithms::MergeUtils; use crate::config::{ModelSoupConfig, SoupMethod}; use crate::error::{MergeError, Result}; use crate::strategies::StrategyUtils; use crate::types::{ MergeInfo, MergeStatistics, MergedModel, Model, ModelReference, QualityMetrics, }; use indexmap::IndexMap; use std::collections::HashMap; use tracing::{debug, info, warn}; /// Model Soup merging strategy implementation pub struct ModelSoupMerger { config: ModelSoupConfig, } impl ModelSoupMerger { /// Create a new Model Soup merger with configuration pub fn new(config: ModelSoupConfig) -> Self { Self { config } } /// Execute Model Soup merging on multiple models pub async fn merge(&self, models: &[Model]) -> Result { if models.len() < 2 { return Err(MergeError::algorithm( "ModelSoup", "At least 2 models required for soup creation", )); } info!("Starting Model Soup creation with {} models", models.len()); let start_time = std::time::Instant::now(); // Step 1: Validate model compatibility StrategyUtils::validate_model_group(models)?; // Step 2: Apply soup construction strategy let (selected_models, final_weights) = match self.config.soup_method { SoupMethod::Uniform => self.uniform_soup_construction(models).await?, SoupMethod::Weighted => self.weighted_soup_construction(models).await?, SoupMethod::Greedy => self.greedy_soup_construction(models).await?, }; // Step 3: Create the soup by averaging selected models let merged_parameters = self.create_soup(&selected_models, &final_weights).await?; // Step 4: Create merged model let mut merged_model = models[0].clone(); merged_model.id = uuid::Uuid::new_v4(); merged_model.name = format!( "ModelSoup_{}_{}", self.config.soup_method.name(), selected_models.len() ); merged_model.parameters = merged_parameters; // Step 5: Compute merge statistics and quality metrics let statistics = self.compute_statistics( &merged_model, models, &selected_models, start_time.elapsed(), ); let quality_metrics = self.compute_quality_metrics(&merged_model, &selected_models, &final_weights)?; let merge_info = MergeInfo { strategy: "ModelSoup".to_string(), source_models: selected_models .iter() .zip(final_weights.iter()) .map(|(model, &weight)| ModelReference { id: model.id, name: model.name.clone(), path: std::path::PathBuf::from(&model.name), weight: Some(weight), }) .collect(), merged_at: chrono::Utc::now(), config: serde_json::to_value(&self.config) .map_err(|e| MergeError::internal(format!("Config serialization: {e}")))?, statistics, }; info!( "Model Soup creation completed in {:?}", start_time.elapsed() ); Ok(MergedModel { model: merged_model, merge_info, quality_metrics, }) } /// Uniform soup construction: equal weights for all models async fn uniform_soup_construction(&self, models: &[Model]) -> Result<(Vec, Vec)> { info!("Creating uniform soup with equal weights"); let selected_count = self.config.max_models.min(models.len()); let selected_models = if selected_count < models.len() { // Select most diverse subset if we need to limit self.select_diverse_subset(models, selected_count)? } else { models.to_vec() }; let uniform_weight = 1.0 / selected_models.len() as f32; let weights = vec![uniform_weight; selected_models.len()]; info!( "Selected {} models with uniform weights: {:.4}", selected_models.len(), uniform_weight ); Ok((selected_models, weights)) } /// Weighted soup construction: weights based on model performance async fn weighted_soup_construction(&self, models: &[Model]) -> Result<(Vec, Vec)> { info!("Creating weighted soup based on model performance"); // Get performance scores (in practice, these would come from validation) let performance_scores = StrategyUtils::compute_performance_scores(models); // Select top-performing models up to max_models let mut model_scores: Vec<(usize, f32)> = models .iter() .enumerate() .map(|(i, model)| { ( i, performance_scores.get(&model.name).copied().unwrap_or(0.5), ) }) .collect(); // Sort by performance (descending) model_scores.sort_by(|a, b| b.1.total_cmp(&a.1)); let selected_count = self.config.max_models.min(models.len()); let selected_indices: Vec = model_scores .iter() .take(selected_count) .map(|(i, _)| *i) .collect(); let selected_models: Vec = selected_indices .iter() .map(|&i| models[i].clone()) .collect(); // Compute softmax weights from performance scores let selected_scores: Vec = model_scores .iter() .take(selected_count) .map(|(_, score)| *score) .collect(); let weights = self.compute_softmax_weights(&selected_scores)?; info!( "Selected {} models with performance-based weights (max: {:.4}, min: {:.4})", selected_models.len(), weights.iter().max_by(|a, b| a.total_cmp(b)).unwrap(), weights.iter().min_by(|a, b| a.total_cmp(b)).unwrap() ); Ok((selected_models, weights)) } /// Greedy soup construction: iteratively add models that improve performance async fn greedy_soup_construction(&self, models: &[Model]) -> Result<(Vec, Vec)> { info!("Creating greedy soup with iterative model selection"); let performance_scores = StrategyUtils::compute_performance_scores(models); // Start with the best-performing model let mut best_idx = 0; let mut best_score = 0.0; for (i, model) in models.iter().enumerate() { let score = performance_scores.get(&model.name).copied().unwrap_or(0.5); if score > best_score { best_score = score; best_idx = i; } } let mut selected_models = vec![models[best_idx].clone()]; let mut current_soup_score = best_score; let mut remaining_indices: Vec = (0..models.len()).filter(|&i| i != best_idx).collect(); info!( "Starting greedy soup with model '{}' (score: {:.4})", models[best_idx].name, best_score ); // Iteratively add models that improve the soup while selected_models.len() < self.config.max_models && !remaining_indices.is_empty() { let mut best_addition_idx = None; let mut best_addition_score = current_soup_score; for (candidate_pos, &candidate_idx) in remaining_indices.iter().enumerate() { // Create temporary soup with candidate model let mut temp_soup_models = selected_models.clone(); temp_soup_models.push(models[candidate_idx].clone()); // Compute temporary soup score let temp_score = self.evaluate_soup_performance(&temp_soup_models).await?; debug!( "Candidate '{}': soup score {:.4} vs current {:.4}", models[candidate_idx].name, temp_score, current_soup_score ); if temp_score > best_addition_score { best_addition_score = temp_score; best_addition_idx = Some(candidate_pos); } } if let Some(addition_pos) = best_addition_idx { let candidate_idx = remaining_indices.remove(addition_pos); selected_models.push(models[candidate_idx].clone()); current_soup_score = best_addition_score; info!( "Added model '{}' to soup, new score: {:.4}", models[candidate_idx].name, current_soup_score ); } else { info!("No more models improve the soup, stopping greedy construction"); break; } } // Use uniform weights for greedy soup (can be enhanced with optimization) let uniform_weight = 1.0 / selected_models.len() as f32; let weights = vec![uniform_weight; selected_models.len()]; info!( "Greedy soup completed with {} models", selected_models.len() ); Ok((selected_models, weights)) } /// Evaluate soup performance using comprehensive model analysis async fn evaluate_soup_performance(&self, soup_models: &[Model]) -> Result { if soup_models.is_empty() { return Ok(0.0); } // 1. Compute ensemble diversity benefits let diversity_score = self.compute_ensemble_diversity(soup_models)?; // 2. Analyze architectural compatibility let compatibility_score = self.compute_architectural_compatibility(soup_models)?; // 3. Evaluate parameter quality across models let quality_score = self.compute_ensemble_quality(soup_models)?; // 4. Consider ensemble size effects let size_factor = self.compute_ensemble_size_factor(soup_models.len()); // 5. Model complementarity analysis let complementarity_score = self.compute_model_complementarity(soup_models)?; // Weighted combination of factors let performance_score = 0.25 * diversity_score + 0.20 * compatibility_score + 0.25 * quality_score + 0.15 * complementarity_score + 0.15 * size_factor; Ok(performance_score.min(1.0).max(0.0)) } /// Compute ensemble diversity score fn compute_ensemble_diversity(&self, models: &[Model]) -> Result { let mut pairwise_distances = Vec::new(); for i in 0..models.len() { for j in (i + 1)..models.len() { let distance = self.compute_model_distance(&models[i], &models[j])?; pairwise_distances.push(distance); } } if pairwise_distances.is_empty() { return Ok(0.0); } let avg_distance = pairwise_distances.iter().sum::() / pairwise_distances.len() as f32; Ok(avg_distance.min(1.0)) } /// Compute distance between two models based on parameter differences fn compute_model_distance(&self, model1: &Model, model2: &Model) -> Result { let mut total_distance = 0.0; let mut param_count = 0; for (param_name, param1) in &model1.parameters { if let Some(param2) = model2.parameters.get(param_name) && param1.data.len() == param2.data.len() { // Compute L2 distance between parameters let l2_dist: f32 = param1 .data .iter() .zip(param2.data.iter()) .map(|(&a, &b)| (a - b).powi(2)) .sum::() .sqrt(); // Normalize by parameter size let normalized_dist = l2_dist / (param1.data.len() as f32).sqrt(); total_distance += normalized_dist; param_count += 1; } } if param_count == 0 { return Ok(0.0); } Ok((total_distance / param_count as f32).min(10.0) / 10.0) // Normalize to [0,1] } /// Compute architectural compatibility score fn compute_architectural_compatibility(&self, models: &[Model]) -> Result { if models.len() <= 1 { return Ok(1.0); } let base_arch = &models[0].architecture; let mut compatibility_scores = Vec::new(); for model in models.iter().skip(1) { let arch_score = self.compute_architecture_compatibility_score(base_arch, &model.architecture); compatibility_scores.push(arch_score); } let avg_compatibility = compatibility_scores.iter().sum::() / compatibility_scores.len() as f32; Ok(avg_compatibility) } /// Compute compatibility between two architectures fn compute_architecture_compatibility_score( &self, arch1: &crate::types::ModelArchitecture, arch2: &crate::types::ModelArchitecture, ) -> f32 { // Architecture type compatibility let type_compatibility = if arch1.arch_type == arch2.arch_type { 1.0 } else { 0.5 // Different types can still be compatible }; // Layer count compatibility let layer_diff = (arch1.num_layers as f32 - arch2.num_layers as f32).abs(); let layer_compatibility = (1.0 - layer_diff / 24.0).max(0.0); // Assume max 24 layers difference // Hidden dimension compatibility let hidden_diff = (arch1.hidden_dim as f32 - arch2.hidden_dim as f32).abs(); let hidden_compatibility = (1.0 - hidden_diff / 2048.0).max(0.0); // Assume max 2048 difference (type_compatibility + layer_compatibility + hidden_compatibility) / 3.0 } /// Compute ensemble quality score fn compute_ensemble_quality(&self, models: &[Model]) -> Result { let mut individual_qualities = Vec::new(); for model in models { let quality = self.compute_individual_model_quality(model); individual_qualities.push(quality); } if individual_qualities.is_empty() { return Ok(0.5); } // Use harmonic mean to penalize low-quality models in ensemble let harmonic_mean = individual_qualities.len() as f32 / individual_qualities .iter() .map(|&q| 1.0 / q.max(0.1)) .sum::(); Ok(harmonic_mean.min(1.0)) } /// Compute quality score for individual model fn compute_individual_model_quality(&self, model: &Model) -> f32 { // Parameter distribution health let param_health = self.compute_parameter_health(model); // Architecture appropriateness let arch_quality = self.compute_architecture_quality(model); // Model size reasonableness let size_quality = self.compute_size_quality(model); (param_health + arch_quality + size_quality) / 3.0 } /// Compute parameter health score fn compute_parameter_health(&self, model: &Model) -> f32 { let mut health_scores = Vec::new(); for (_param_name, param) in &model.parameters { if param.data.is_empty() { continue; } let (mean, std_dev, skewness, kurtosis) = MergeUtils::compute_moments(¶m.data); // Check for reasonable statistics let mean_score = 1.0 - mean.abs().min(1.0); // Prefer small mean let std_score = if std_dev > 1e-6 && std_dev < 2.0 { 1.0 } else { 0.5 }; // Reasonable std let skew_score = 1.0 - skewness.abs().min(2.0) / 2.0; // Not too skewed let kurt_score = 1.0 - (kurtosis - 3.0).abs().min(3.0) / 3.0; // Near normal kurtosis let param_score = (mean_score + std_score + skew_score + kurt_score) / 4.0; health_scores.push(param_score); } if health_scores.is_empty() { return 0.5; } health_scores.iter().sum::() / health_scores.len() as f32 } /// Compute architecture quality score fn compute_architecture_quality(&self, model: &Model) -> f32 { let arch = &model.architecture; // Prefer well-known architectures let type_score = match arch.arch_type.as_str() { "transformer" => 1.0, "cnn" => 0.9, "rnn" | "lstm" => 0.8, "feedforward" => 0.7, _ => 0.5, }; // Reasonable layer count let layer_score = if arch.num_layers >= 4 && arch.num_layers <= 32 { 1.0 } else { 0.7 }; // Standard hidden dimensions let hidden_score = match arch.hidden_dim { 256 | 384 | 512 | 768 | 1024 | 1536 | 2048 => 1.0, _ => 0.8, }; (type_score + layer_score + hidden_score) / 3.0 } /// Compute size quality score fn compute_size_quality(&self, model: &Model) -> f32 { let param_count = model.parameter_count(); // Sweet spot for model size if (1_000_000..=1_000_000_000).contains(¶m_count) { 1.0 } else if param_count < 1_000_000 { 0.7 // Too small } else { 0.8 // Very large but acceptable } } /// Compute ensemble size factor (diminishing returns) fn compute_ensemble_size_factor(&self, ensemble_size: usize) -> f32 { match ensemble_size { 0 => 0.0, 1 => 0.3, // Single model 2..=3 => 0.8, // Good ensemble size 4..=6 => 1.0, // Optimal ensemble size 7..=10 => 0.9, // Still good but diminishing returns _ => 0.7, // Too many models may hurt performance } } /// Compute model complementarity (how well models complement each other) fn compute_model_complementarity(&self, models: &[Model]) -> Result { if models.len() <= 1 { return Ok(0.0); } // Analyze parameter variance patterns across models let mut complementarity_scores = Vec::new(); // Sample some parameters to analyze complementarity if let Some(first_model) = models.first() { for (param_name, _) in first_model.parameters.iter().take(5) { // Sample 5 parameters let complementarity = self.compute_parameter_complementarity(models, param_name)?; complementarity_scores.push(complementarity); } } if complementarity_scores.is_empty() { return Ok(0.5); } Ok(complementarity_scores.iter().sum::() / complementarity_scores.len() as f32) } /// Compute complementarity for a specific parameter across models fn compute_parameter_complementarity(&self, models: &[Model], param_name: &str) -> Result { let mut param_values = Vec::new(); for model in models { if let Some(param) = model.parameters.get(param_name) { // Take mean of parameter for simplicity let param_mean = param.data.iter().sum::() / param.data.len() as f32; param_values.push(param_mean); } } if param_values.len() <= 1 { return Ok(0.0); } // Compute variance as measure of complementarity let mean = param_values.iter().sum::() / param_values.len() as f32; let variance = param_values .iter() .map(|&x| (x - mean).powi(2)) .sum::() / param_values.len() as f32; // Normalize variance to [0,1] scale let normalized_variance = (variance.sqrt()).min(1.0); Ok(normalized_variance) } /// Select diverse subset of models fn select_diverse_subset(&self, models: &[Model], max_count: usize) -> Result> { if models.len() <= max_count { return Ok(models.to_vec()); } let all_indices: Vec = (0..models.len()).collect(); let selected_indices = StrategyUtils::select_representatives(models, &all_indices, max_count)?; Ok(selected_indices .into_iter() .map(|i| models[i].clone()) .collect()) } /// Compute softmax weights from scores fn compute_softmax_weights(&self, scores: &[f32]) -> Result> { if scores.is_empty() { return Ok(vec![]); } // Apply softmax with temperature to avoid extreme weights let temperature = 2.0; let max_score = scores.iter().max_by(|a, b| a.total_cmp(b)).unwrap(); let exp_scores: Vec = scores .iter() .map(|&score| ((score - max_score) / temperature).exp()) .collect(); let sum_exp: f32 = exp_scores.iter().sum(); if sum_exp == 0.0 { // Fallback to uniform weights let uniform_weight = 1.0 / scores.len() as f32; return Ok(vec![uniform_weight; scores.len()]); } Ok(exp_scores .iter() .map(|&exp_score| exp_score / sum_exp) .collect()) } /// Create the final soup by averaging selected models with weights async fn create_soup( &self, models: &[Model], weights: &[f32], ) -> Result> { info!("Creating soup from {} models", models.len()); if models.is_empty() { return Err(MergeError::algorithm("ModelSoup", "No models to merge")); } if models.len() != weights.len() { return Err(MergeError::parameter_mismatch( format!("{} models", models.len()), format!("{} weights", weights.len()), )); } let base_model = &models[0]; let mut soup_parameters = IndexMap::new(); for (param_name, base_param) in &base_model.parameters { debug!("Creating soup for parameter: {}", param_name); // Collect parameter values from all models let mut param_values = Vec::new(); for model in models { if let Some(param) = model.parameters.get(param_name) { param_values.push(param.data.as_slice()); } else { warn!( "Parameter '{}' missing in model '{}'", param_name, model.name ); param_values.push(base_param.data.as_slice()); // Use base as fallback } } // Compute weighted average let soup_data = MergeUtils::weighted_average_parameters(¶m_values, weights)?; // Create soup parameter tensor let mut soup_param = base_param.clone(); soup_param.data = soup_data; soup_parameters.insert(param_name.clone(), soup_param); } Ok(soup_parameters) } /// Compute merge statistics fn compute_statistics( &self, merged_model: &Model, _original_models: &[Model], _selected_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, // Model soup doesn't have explicit conflicts parameters_dropped: 0, // Model soup doesn't drop parameters duration_ms: duration.as_millis() as u64, memory_usage_mb, gpu_memory_mb: None, } } /// Compute quality metrics for the soup fn compute_quality_metrics( &self, soup_model: &Model, selected_models: &[Model], weights: &[f32], ) -> Result { let mut consistency_scores = Vec::new(); // Compute consistency with each selected model for model in selected_models { let similarity = MergeUtils::compute_similarity(soup_model, model)?; consistency_scores.push(similarity); } let consistency_score = if consistency_scores.is_empty() { 0.0 } else { consistency_scores.iter().sum::() / consistency_scores.len() as f32 }; // Compute diversity metrics let selected_indices: Vec = (0..selected_models.len()).collect(); let diversity_score = StrategyUtils::compute_group_diversity(selected_models, &selected_indices)?; // Compute weight distribution metrics let weight_entropy = self.compute_weight_entropy(weights)?; let weight_balance = 1.0 - (weights .iter() .map(|&w| (w - 1.0 / weights.len() as f32).abs()) .sum::() / 2.0); // Compute complexity score based on parameter distribution let all_params: Vec = soup_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("diversity_score".to_string(), diversity_score); quality_indicators.insert("weight_entropy".to_string(), weight_entropy); quality_indicators.insert("weight_balance".to_string(), weight_balance); quality_indicators.insert("parameter_diversity".to_string(), std_dev); quality_indicators.insert("average_similarity".to_string(), consistency_score); quality_indicators.insert( "soup_effectiveness".to_string(), diversity_score * consistency_score, ); let mut validation_results = HashMap::new(); validation_results.insert("diversity_achieved".to_string(), diversity_score > 0.1); validation_results.insert( "consistency_maintained".to_string(), consistency_score > 0.7, ); validation_results.insert("balanced_weights".to_string(), weight_balance > 0.5); validation_results.insert("good_coverage".to_string(), selected_models.len() > 1); let mut performance_predictions = HashMap::new(); performance_predictions.insert("expected_accuracy".to_string(), consistency_score * 0.95); performance_predictions.insert("ensemble_benefit".to_string(), diversity_score * 0.8); performance_predictions.insert( "generalization".to_string(), (diversity_score + consistency_score) * 0.5, ); performance_predictions.insert("robustness".to_string(), weight_balance * 0.9); Ok(QualityMetrics { consistency_score, complexity_score, quality_indicators, validation_results, performance_predictions, }) } /// Compute entropy of weight distribution fn compute_weight_entropy(&self, weights: &[f32]) -> Result { if weights.is_empty() { return Ok(0.0); } let mut entropy = 0.0; for &weight in weights { if weight > 0.0 { entropy -= weight * weight.ln(); } } Ok(entropy) } } impl SoupMethod { /// Get the name of the soup method pub fn name(&self) -> &'static str { match self { Self::Uniform => "Uniform", Self::Weighted => "Weighted", Self::Greedy => "Greedy", } } } /// Convenience function for Model Soup merging pub async fn merge_models(models: &[Model], config: &ModelSoupConfig) -> Result { let merger = ModelSoupMerger::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_uniform_soup() -> 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![2.0, 3.0, 4.0])]); let model3 = create_test_model("model3", vec![("weight".to_string(), vec![3.0, 4.0, 5.0])]); let config = ModelSoupConfig { soup_method: SoupMethod::Uniform, max_models: 10, ..ModelSoupConfig::default() }; let merger = ModelSoupMerger::new(config); let result = merger.merge(&[model1, model2, model3]).await?; assert!(result.model.parameters.contains_key("weight")); assert_eq!(result.merge_info.strategy, "ModelSoup"); let soup_weight = result.model.get_parameter("weight").unwrap(); // Should be average: (1+2+3)/3, (2+3+4)/3, (3+4+5)/3 = [2.0, 3.0, 4.0] assert_eq!(soup_weight.data, vec![2.0, 3.0, 4.0]); Ok(()) } #[tokio::test] async fn test_weighted_soup() -> Result<()> { let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 1.0])]); let model2 = create_test_model("model2", vec![("weight".to_string(), vec![3.0, 3.0])]); let config = ModelSoupConfig { soup_method: SoupMethod::Weighted, max_models: 10, ..ModelSoupConfig::default() }; let merger = ModelSoupMerger::new(config); let result = merger.merge(&[model1, model2]).await?; assert!(result.model.parameters.contains_key("weight")); let soup_weight = result.model.get_parameter("weight").unwrap(); // Should be weighted average (exact values depend on performance scoring) assert!(soup_weight.data[0] > 1.0 && soup_weight.data[0] < 3.0); assert!(soup_weight.data[1] > 1.0 && soup_weight.data[1] < 3.0); Ok(()) } #[tokio::test] async fn test_greedy_soup() -> 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![2.0, 3.0])]); let model3 = create_test_model("model3", vec![("weight".to_string(), vec![3.0, 4.0])]); let config = ModelSoupConfig { soup_method: SoupMethod::Greedy, max_models: 2, // Limit to test selection ..ModelSoupConfig::default() }; let merger = ModelSoupMerger::new(config); let result = merger.merge(&[model1, model2, model3]).await?; assert!(result.model.parameters.contains_key("weight")); assert_eq!(result.merge_info.source_models.len(), 2); // Should select 2 models Ok(()) } #[test] fn test_softmax_weights_computation() -> Result<()> { let config = ModelSoupConfig::default(); let merger = ModelSoupMerger::new(config); let scores = vec![0.8, 0.6, 0.9, 0.7]; let weights = merger.compute_softmax_weights(&scores)?; assert_eq!(weights.len(), 4); // Weights should sum to 1.0 let weight_sum: f32 = weights.iter().sum(); assert!((weight_sum - 1.0).abs() < 1e-6); // Higher scores should get higher weights let max_score_idx = scores .iter() .enumerate() .max_by(|(_, a), (_, b)| a.total_cmp(b)) .unwrap() .0; let max_weight_idx = weights .iter() .enumerate() .max_by(|(_, a), (_, b)| a.total_cmp(b)) .unwrap() .0; assert_eq!(max_score_idx, max_weight_idx); Ok(()) } #[test] fn test_diverse_subset_selection() -> Result<()> { let models = (0..5) .map(|i| { create_test_model( &format!("model{}", i), vec![("weight".to_string(), vec![i as f32, (i * 2) as f32])], ) }) .collect::>(); let config = ModelSoupConfig::default(); let merger = ModelSoupMerger::new(config); let subset = merger.select_diverse_subset(&models, 3)?; assert_eq!(subset.len(), 3); // Should select diverse models (can't easily test exact selection without knowing algorithm details) Ok(()) } #[test] fn test_weight_entropy() -> Result<()> { let config = ModelSoupConfig::default(); let merger = ModelSoupMerger::new(config); // Uniform weights should have high entropy let uniform_weights = vec![0.25, 0.25, 0.25, 0.25]; let uniform_entropy = merger.compute_weight_entropy(&uniform_weights)?; // Skewed weights should have lower entropy let skewed_weights = vec![0.7, 0.1, 0.1, 0.1]; let skewed_entropy = merger.compute_weight_entropy(&skewed_weights)?; assert!(uniform_entropy > skewed_entropy); Ok(()) } #[tokio::test] async fn test_soup_performance_evaluation() -> Result<()> { let model1 = create_test_model( "diverse", vec![ ("weight".to_string(), vec![1.0, 5.0, 10.0]), // High diversity ], ); let model2 = create_test_model( "uniform", vec![ ("weight".to_string(), vec![2.0, 2.0, 2.0]), // Low diversity ], ); let config = ModelSoupConfig::default(); let merger = ModelSoupMerger::new(config); let diverse_score = merger.evaluate_soup_performance(&[model1.clone()]).await?; let uniform_score = merger.evaluate_soup_performance(&[model2.clone()]).await?; let combined_score = merger.evaluate_soup_performance(&[model1, model2]).await?; assert!(diverse_score > uniform_score); // Combined soup should potentially have different score assert!(combined_score >= 0.0 && combined_score <= 1.0); Ok(()) } #[tokio::test] async fn test_max_models_limit() -> Result<()> { let models = (0..5) .map(|i| { create_test_model( &format!("model{}", i), vec![("weight".to_string(), vec![i as f32])], ) }) .collect::>(); let config = ModelSoupConfig { soup_method: SoupMethod::Uniform, max_models: 3, // Limit to 3 models ..ModelSoupConfig::default() }; let merger = ModelSoupMerger::new(config); let result = merger.merge(&models).await?; // Should only use 3 models assert_eq!(result.merge_info.source_models.len(), 3); Ok(()) } }