//! Integration tests for RTX Model Merging use rtx_model_merging::*; use std::collections::HashMap; use tempfile::TempDir; use uuid::Uuid; fn create_test_model(name: &str, param_values: Vec) -> types::Model { let arch = types::ModelArchitecture { arch_type: "test".to_string(), num_layers: 1, hidden_dim: param_values.len(), params: HashMap::new(), }; let mut model = types::Model::new(name.to_string(), arch); let param = types::ParameterTensor::new( "weight".to_string(), vec![param_values.len()], types::DataType::Float32, param_values, ); model.add_parameter(param); model } #[tokio::test] async fn test_end_to_end_ties_merge() -> Result<()> { let model1 = create_test_model("model1", vec![1.0, 2.0, 3.0]); let model2 = create_test_model("model2", vec![1.5, 2.5, 3.5]); let model3 = create_test_model("model3", vec![0.5, 1.5, 2.5]); let config = config::TiesConfig { sign_threshold: 0.6, magnitude_threshold: 0.1, density: 0.8, enable_sign_voting: true, voting_threshold: 0.5, rescale_method: config::RescaleMethod::Magnitude, }; let merged = algorithms::ties::merge_models(&[model1, model2, model3], &config).await?; assert_eq!(merged.merge_info.strategy, "TIES"); assert_eq!(merged.merge_info.source_models.len(), 3); assert!(merged.model.parameters.contains_key("weight")); let weight = merged.model.get_parameter("weight").unwrap(); assert_eq!(weight.data.len(), 3); Ok(()) } #[tokio::test] async fn test_end_to_end_model_soup() -> Result<()> { let model1 = create_test_model("model1", vec![2.0, 3.0]); let model2 = create_test_model("model2", vec![4.0, 5.0]); let model3 = create_test_model("model3", vec![6.0, 7.0]); let config = config::ModelSoupConfig { weights: vec![], soup_method: config::SoupMethod::Uniform, max_models: 5, greedy_soup: false, validation_metric: "accuracy".to_string(), }; let merged = strategies::model_soup::merge_models(&[model1, model2, model3], &config).await?; assert_eq!(merged.merge_info.strategy, "ModelSoup"); assert!(merged.model.parameters.contains_key("weight")); let weight = merged.model.get_parameter("weight").unwrap(); // Should be average: (2+4+6)/3 = 4, (3+5+7)/3 = 5 assert_eq!(weight.data, vec![4.0, 5.0]); Ok(()) } #[tokio::test] async fn test_model_merger_workflow() -> Result<()> { let model1 = create_test_model("base", vec![1.0, 2.0, 3.0]); let model2 = create_test_model("task1", vec![1.2, 2.3, 3.1]); // Test actual merging with SLERP let slerp_config = config::SlerpConfig { t: 0.5, use_quaternions: false, normalization: config::NormalizationMethod::L2, adaptive_t: false, parameter_weights: None, }; let result = algorithms::slerp::merge_models(&[model1, model2], &slerp_config).await?; assert!(result.model.parameters.contains_key("weight")); assert!(!result.quality_metrics.quality_indicators.is_empty()); assert!(result.quality_metrics.consistency_score >= 0.0); Ok(()) } #[tokio::test] async fn test_conflict_resolution() -> Result<()> { let validator = conflict_resolution::MergeValidator::new(); // Create conflicting parameter vectors let vec1 = vec![1.0, -2.0, 3.0]; let vec2 = vec![-1.0, 2.0, -3.0]; let vec3 = vec![5.0, 10.0, 15.0]; let param_vectors = vec![ vec1.as_slice(), // Positive, negative, positive vec2.as_slice(), // Negative, positive, negative (sign conflicts) vec3.as_slice(), // Large magnitudes ]; let conflicts = validator.detect_conflicts("test_param", ¶m_vectors)?; assert!(!conflicts.is_empty()); // Test conflict resolution let resolved = validator.resolve_conflicts(&conflicts, ¶m_vectors)?; assert_eq!(resolved.len(), 3); // Values should be resolved (exact values depend on resolution strategy) for &value in &resolved { assert!(value.is_finite()); } Ok(()) } #[tokio::test] async fn test_performance_evaluation() -> Result<()> { let evaluator = evaluation::PerformanceEvaluator::new()?; let model = create_test_model("test", vec![1.0, 2.0, 3.0, 4.0]); let merged_model = types::MergedModel { model, merge_info: types::MergeInfo { strategy: "Test".to_string(), source_models: vec![], merged_at: chrono::Utc::now(), config: serde_json::json!({}), statistics: types::MergeStatistics::default(), }, quality_metrics: types::QualityMetrics::default(), }; let results = evaluator.evaluate(&merged_model).await?; assert!(results.overall_score >= 0.0 && results.overall_score <= 1.0); assert!(!results.metrics.is_empty()); assert!(results.metrics.contains_key("accuracy")); assert!(results.comparisons.is_empty()); // No source models Ok(()) } #[tokio::test] async fn test_merge_planning() -> Result<()> { let planner = planning::MergePlanner::new(); let model1 = create_test_model("model1", vec![1.0, 2.0]); let model2 = create_test_model("model2", vec![3.0, 4.0]); // Test compatibility analysis let report = planner.analyze_compatibility(&[model1.clone(), model2.clone()])?; assert!(report.overall_compatible); assert_eq!(report.pairwise_compatibility.len(), 1); // Test execution planning let strategy = config::MergeStrategy::Ties(config::TiesConfig::default()); let constraints = planning::ResourceConstraints::default(); let plan = planner.plan_merge_execution(&[model1, model2], &strategy, &constraints)?; assert!(!plan.phases.is_empty()); assert!(plan.estimated_duration_ms > 0); assert!(!plan.checkpoints.is_empty()); Ok(()) } #[tokio::test] async fn test_configuration_system() -> Result<()> { let temp_dir = TempDir::new().unwrap(); let config_path = temp_dir.path().join("test_config.json"); // Create and save configuration let original_config = config::MergeConfig::default(); config::save_config(&original_config, &config_path)?; // Load configuration let loaded_config = config::load_config(&config_path)?; assert_eq!( original_config.loader.use_memory_mapping, loaded_config.loader.use_memory_mapping ); assert_eq!( original_config.validation.validate_architecture, loaded_config.validation.validate_architecture ); Ok(()) } #[tokio::test] async fn test_different_merge_strategies() -> Result<()> { let models = vec![ create_test_model("model1", vec![1.0, 2.0]), create_test_model("model2", vec![2.0, 3.0]), ]; // Test DARE let dare_config = config::DareConfig { drop_probability: 0.1, rescale_factor: 1.0, seed: Some(42), adaptive_dropping: false, importance_threshold: 0.01, }; let dare_result = algorithms::dare::merge_models(&models, &dare_config).await?; assert_eq!(dare_result.merge_info.strategy, "DARE"); // Test Task Arithmetic let task_config = config::TaskArithmeticConfig { scaling_factors: vec![1.5], preserve_signs: true, magnitude_clip: None, selective_arithmetic: false, selection_criteria: vec![], }; let task_result = algorithms::task_arithmetic::merge_models(&models, &task_config).await?; assert_eq!(task_result.merge_info.strategy, "TaskArithmetic"); // Test Fisher let fisher_config = config::FisherConfig { fisher_method: config::FisherMethod::Diagonal, regularization: 1e-6, diagonal_fisher: true, num_samples: 100, empirical_fisher: true, }; let fisher_result = algorithms::fisher::merge_models(&models, &fisher_config).await?; assert_eq!(fisher_result.merge_info.strategy, "Fisher"); Ok(()) } #[tokio::test] async fn test_model_validation() -> Result<()> { let validator = conflict_resolution::MergeValidator::new(); let valid_model = create_test_model("valid", vec![1.0, 2.0, 3.0]); let models = vec![valid_model]; // Test compatibility validation let result = validator.validate_compatibility(&models); assert!(result.is_ok()); // Test merged model validation let merged_model = types::MergedModel { model: models[0].clone(), merge_info: types::MergeInfo { strategy: "Test".to_string(), source_models: vec![], merged_at: chrono::Utc::now(), config: serde_json::json!({}), statistics: types::MergeStatistics::default(), }, quality_metrics: types::QualityMetrics::default(), }; let validation_result = validator.validate_merged_model(&merged_model); assert!(validation_result.is_ok()); Ok(()) } #[tokio::test] async fn test_large_model_merge() -> Result<()> { // Create larger models to test scalability let large_params1: Vec = (0..1000).map(|x| (x as f32) * 0.01).collect(); let large_params2: Vec = (0..1000).map(|x| (x as f32) * 0.011).collect(); let model1 = create_test_model("large1", large_params1); let model2 = create_test_model("large2", large_params2); let config = config::SlerpConfig { t: 0.5, normalization: config::NormalizationMethod::L2, ..config::SlerpConfig::default() }; let result = algorithms::slerp::merge_models(&[model1, model2], &config).await?; assert_eq!( result.model.get_parameter("weight").unwrap().data.len(), 1000 ); assert!(result.quality_metrics.consistency_score > 0.0); Ok(()) } #[tokio::test] async fn test_error_handling() -> Result<()> { // Test merging with incompatible models let model1 = create_test_model("model1", vec![1.0, 2.0]); let mut model2 = create_test_model("model2", vec![1.0, 2.0, 3.0]); // Different size model2.architecture.hidden_dim = 3; // Make architecture incompatible let config = config::TiesConfig::default(); let result = algorithms::ties::merge_models(&[model1, model2], &config).await; // Should fail due to incompatibility assert!(result.is_err()); Ok(()) } #[tokio::test] async fn test_quality_metrics() -> Result<()> { let model1 = create_test_model("model1", vec![1.0, 2.0, 3.0]); let model2 = create_test_model("model2", vec![1.1, 2.1, 3.1]); let config = config::TiesConfig::default(); let result = algorithms::ties::merge_models(&[model1, model2], &config).await?; let metrics = &result.quality_metrics; assert!(metrics.consistency_score >= 0.0 && metrics.consistency_score <= 1.0); assert!(metrics.complexity_score >= 0.0); assert!(!metrics.quality_indicators.is_empty()); assert!(!metrics.validation_results.is_empty()); assert!(!metrics.performance_predictions.is_empty()); Ok(()) } #[test] fn test_utility_functions() -> Result<()> { use rtx_model_merging::algorithms::MergeUtils; // Test cosine similarity let vec1 = vec![1.0, 2.0, 3.0]; let vec2 = vec![1.0, 2.0, 3.0]; let similarity = MergeUtils::cosine_similarity(&vec1, &vec2)?; assert!((similarity - 1.0).abs() < 1e-6); // Test parameter averaging let param1 = vec![1.0, 2.0]; let param2 = vec![3.0, 4.0]; let params = vec![param1.as_slice(), param2.as_slice()]; let avg = MergeUtils::average_parameters(¶ms)?; assert_eq!(avg, vec![2.0, 3.0]); // Test weighted averaging let weights = vec![0.3, 0.7]; let weighted_avg = MergeUtils::weighted_average_parameters(¶ms, &weights)?; assert_eq!(weighted_avg, vec![2.4, 3.4]); Ok(()) } #[tokio::test] async fn test_memory_efficiency() -> Result<()> { // Test that merging doesn't create excessive memory overhead let models: Vec = (0..5) .map(|i| create_test_model(&format!("model{}", i), vec![i as f32; 100])) .collect(); let config = config::ModelSoupConfig { weights: vec![], soup_method: config::SoupMethod::Uniform, max_models: 5, greedy_soup: false, validation_metric: "accuracy".to_string(), }; let initial_memory: usize = models.iter().map(|m| m.memory_size()).sum(); let result = strategies::model_soup::merge_models(&models, &config).await?; let final_memory = result.model.memory_size(); // Merged model should not be significantly larger than individual models assert!(final_memory <= initial_memory); Ok(()) }