392 lines
12 KiB
Rust
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", ¶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<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(¶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<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(())
|
|
}
|