Files
rustytorch/crates/training/rtx-model-merging/tests/integration_tests.rs
T
2026-03-04 00:08:42 +00:00

392 lines
12 KiB
Rust

//! 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<f32>) -> 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", &param_vectors)?;
assert!(!conflicts.is_empty());
// Test conflict resolution
let resolved = validator.resolve_conflicts(&conflicts, &param_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<f32> = (0..1000).map(|x| (x as f32) * 0.01).collect();
let large_params2: Vec<f32> = (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(&params)?;
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(&params, &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<types::Model> = (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(())
}