//! Conflict resolution system for model parameter merging use crate::algorithms::MergeUtils; use crate::config::ValidationConfig; use crate::error::{MergeError, Result}; use crate::types::{ConflictResolution, ConflictType, ParameterConflict}; use tracing::warn; /// Model merge validator and conflict resolver pub struct MergeValidator { config: ValidationConfig, } impl MergeValidator { /// Create a new merge validator with default configuration pub fn new() -> Self { Self::with_config(ValidationConfig::default()) } /// Create a new merge validator with custom configuration pub fn with_config(config: ValidationConfig) -> Self { Self { config } } /// Validate model compatibility for merging pub fn validate_compatibility(&self, models: &[crate::types::Model]) -> Result<()> { if !self.config.validate_architecture { return Ok(()); } if models.len() < 2 { return Ok(()); } let base_model = &models[0]; for (_i, model) in models.iter().enumerate().skip(1) { if !base_model.is_compatible_with(model) { return Err(MergeError::architecture_incompatible( base_model.name.clone(), model.name.clone(), )); } } Ok(()) } /// Validate merged model integrity pub fn validate_merged_model(&self, merged_model: &crate::types::MergedModel) -> Result<()> { let mut validation_errors = Vec::new(); // Numerical stability validation if self.config.validate_numerical_stability && let Err(e) = self.validate_numerical_stability(&merged_model.model) { validation_errors.push(format!("Numerical stability: {e}")); } // Parameter validation if self.config.validate_parameters && let Err(e) = self.validate_parameters(&merged_model.model) { validation_errors.push(format!("Parameters: {e}")); } if validation_errors.is_empty() { Ok(()) } else { Err(MergeError::validation(validation_errors)) } } /// Detect conflicts in parameter merging pub fn detect_conflicts( &self, param_name: &str, param_vectors: &[&[f32]], ) -> Result> { let mut conflicts = Vec::new(); if param_vectors.len() < 2 { return Ok(conflicts); } // Check for value differences let value_conflicts = self.detect_value_conflicts(param_name, param_vectors)?; conflicts.extend(value_conflicts); // Check for sign conflicts let sign_conflicts = self.detect_sign_conflicts(param_name, param_vectors)?; conflicts.extend(sign_conflicts); // Check for magnitude differences let magnitude_conflicts = self.detect_magnitude_conflicts(param_name, param_vectors)?; conflicts.extend(magnitude_conflicts); Ok(conflicts) } /// Resolve parameter conflicts using configured strategy pub fn resolve_conflicts( &self, conflicts: &[ParameterConflict], param_vectors: &[&[f32]], ) -> Result> { if param_vectors.is_empty() { return Ok(vec![]); } if conflicts.is_empty() { // No conflicts, use simple averaging return MergeUtils::average_parameters(param_vectors); } let param_len = param_vectors[0].len(); let mut resolved = vec![0.0; param_len]; for i in 0..param_len { let values: Vec = param_vectors.iter().map(|v| v[i]).collect(); let element_conflicts: Vec<&ParameterConflict> = conflicts .iter() .filter(|c| self.affects_element(c, i)) .collect(); resolved[i] = if element_conflicts.is_empty() { // Simple average for non-conflicting elements values.iter().sum::() / values.len() as f32 } else { // Apply conflict resolution self.resolve_element_conflict(&values, &element_conflicts)? }; } Ok(resolved) } fn validate_numerical_stability(&self, model: &crate::types::Model) -> Result<()> { for (param_name, param) in &model.parameters { for &value in ¶m.data { if !value.is_finite() { return Err(MergeError::numerical(format!( "Non-finite value in parameter {param_name}: {value}" ))); } if value.abs() > 1e6 { warn!("Large parameter value in {}: {}", param_name, value); } } // Check for extreme variance let (_, std_dev, _, _) = MergeUtils::compute_moments(¶m.data); if std_dev > self.config.max_parameter_deviation { return Err(MergeError::numerical(format!( "Excessive parameter deviation in {param_name}: {std_dev}" ))); } } Ok(()) } fn validate_parameters(&self, model: &crate::types::Model) -> Result<()> { if model.parameters.is_empty() { return Err(MergeError::validation(vec![ "No parameters found".to_string(), ])); } for (param_name, param) in &model.parameters { if param.data.is_empty() { return Err(MergeError::validation(vec![format!( "Empty parameter: {}", param_name )])); } if param.numel() == 0 { return Err(MergeError::validation(vec![format!( "Zero-size parameter: {}", param_name )])); } } Ok(()) } fn detect_value_conflicts( &self, param_name: &str, param_vectors: &[&[f32]], ) -> Result> { let mut conflicts = Vec::new(); let param_len = param_vectors[0].len(); for i in 0..param_len { let values: Vec = param_vectors.iter().map(|v| v[i]).collect(); let (_, std_dev, _, _) = MergeUtils::compute_moments(&values); if std_dev > self.config.numerical_tolerance * 10.0 { conflicts.push(ParameterConflict { parameter_name: param_name.to_string(), models: vec![], // Would be filled with actual model IDs severity: (std_dev / (std_dev + 1.0)).min(1.0), conflict_type: ConflictType::ValueDifference, suggested_resolution: ConflictResolution::WeightedAverage, }); } } Ok(conflicts) } fn detect_sign_conflicts( &self, param_name: &str, param_vectors: &[&[f32]], ) -> Result> { let mut conflicts = Vec::new(); let param_len = param_vectors[0].len(); for i in 0..param_len { let values: Vec = param_vectors.iter().map(|v| v[i]).collect(); let positive_count = values.iter().filter(|&&x| x > 0.0).count(); let negative_count = values.iter().filter(|&&x| x < 0.0).count(); let total_nonzero = positive_count + negative_count; if total_nonzero > 0 && positive_count > 0 && negative_count > 0 { let conflict_ratio = (positive_count.min(negative_count) as f32) / (total_nonzero as f32); if conflict_ratio > 0.3 { conflicts.push(ParameterConflict { parameter_name: param_name.to_string(), models: vec![], severity: conflict_ratio, conflict_type: ConflictType::SignConflict, suggested_resolution: ConflictResolution::MajorityVote, }); } } } Ok(conflicts) } fn detect_magnitude_conflicts( &self, param_name: &str, param_vectors: &[&[f32]], ) -> Result> { let mut conflicts = Vec::new(); let param_len = param_vectors[0].len(); for i in 0..param_len { let values: Vec = param_vectors.iter().map(|v| v[i].abs()).collect(); let max_val = values.iter().max_by(|a, b| a.total_cmp(b)).unwrap(); let min_val = values.iter().min_by(|a, b| a.total_cmp(b)).unwrap(); if *max_val > 0.0 && *max_val / *min_val > 10.0 { conflicts.push(ParameterConflict { parameter_name: param_name.to_string(), models: vec![], severity: (*max_val / (*max_val + *min_val)).min(1.0), conflict_type: ConflictType::MagnitudeDifference, suggested_resolution: ConflictResolution::Median, }); } } Ok(conflicts) } fn affects_element(&self, _conflict: &ParameterConflict, _element_index: usize) -> bool { // Simplified implementation - in practice would check element-specific conflicts true } fn resolve_element_conflict( &self, values: &[f32], conflicts: &[&ParameterConflict], ) -> Result { if conflicts.is_empty() { return Ok(values.iter().sum::() / values.len() as f32); } // Use the resolution strategy from the most severe conflict let primary_conflict = conflicts .iter() .max_by(|a, b| a.severity.total_cmp(&b.severity)) .unwrap(); match primary_conflict.suggested_resolution { ConflictResolution::Average => Ok(values.iter().sum::() / values.len() as f32), ConflictResolution::WeightedAverage => { // Use inverse severity as weights let weights: Vec = values.iter().map(|_| 1.0).collect(); // Simplified MergeUtils::weighted_average_parameters(&[values], &weights).map(|v| v[0]) } ConflictResolution::MajorityVote => self.resolve_by_majority_vote(values), ConflictResolution::Median => Ok(self.compute_median(values)), ConflictResolution::BestModel => { // Use first value as "best" (simplified) Ok(values[0]) } ConflictResolution::Drop => Ok(0.0), ConflictResolution::Custom(_) => Ok(values.iter().sum::() / values.len() as f32), } } fn resolve_by_majority_vote(&self, values: &[f32]) -> Result { if values.is_empty() { return Ok(0.0); } // Sign-based majority vote let positive_values: Vec = values.iter().filter(|&&x| x > 0.0).copied().collect(); let negative_values: Vec = values.iter().filter(|&&x| x < 0.0).copied().collect(); if positive_values.len() > negative_values.len() { Ok(positive_values.iter().sum::() / positive_values.len() as f32) } else if negative_values.len() > positive_values.len() { Ok(negative_values.iter().sum::() / negative_values.len() as f32) } else { // Tie - use simple average Ok(values.iter().sum::() / values.len() as f32) } } fn compute_median(&self, values: &[f32]) -> f32 { if values.is_empty() { return 0.0; } let mut sorted_values = values.to_vec(); sorted_values.sort_by(f32::total_cmp); let len = sorted_values.len(); if len.is_multiple_of(2) { f32::midpoint(sorted_values[len / 2 - 1], sorted_values[len / 2]) } else { sorted_values[len / 2] } } } impl Default for MergeValidator { fn default() -> Self { Self::new() } } #[cfg(test)] mod tests { use super::*; #[test] fn test_value_conflict_detection() -> Result<()> { let validator = MergeValidator::new(); // High variance values should trigger conflict let vec1 = vec![1.0, 2.0, 3.0]; let vec2 = vec![10.0, 20.0, 30.0]; let param_vectors = vec![ vec1.as_slice(), vec2.as_slice(), // Very different values ]; let conflicts = validator.detect_value_conflicts("test_param", ¶m_vectors)?; assert!(!conflicts.is_empty()); assert_eq!(conflicts[0].conflict_type, ConflictType::ValueDifference); Ok(()) } #[test] fn test_sign_conflict_detection() -> Result<()> { let validator = MergeValidator::new(); let vec1 = vec![1.0, 2.0, 3.0]; let vec2 = vec![-1.0, -2.0, -3.0]; let param_vectors = vec![ vec1.as_slice(), vec2.as_slice(), // Opposite signs ]; let conflicts = validator.detect_sign_conflicts("test_param", ¶m_vectors)?; assert!(!conflicts.is_empty()); assert_eq!(conflicts[0].conflict_type, ConflictType::SignConflict); Ok(()) } #[test] fn test_magnitude_conflict_detection() -> Result<()> { let validator = MergeValidator::new(); let vec1 = vec![0.1, 0.1, 0.1]; let vec2 = vec![10.0, 10.0, 10.0]; let param_vectors = vec![ vec1.as_slice(), vec2.as_slice(), // 100x magnitude difference ]; let conflicts = validator.detect_magnitude_conflicts("test_param", ¶m_vectors)?; assert!(!conflicts.is_empty()); assert_eq!( conflicts[0].conflict_type, ConflictType::MagnitudeDifference ); Ok(()) } #[test] fn test_median_computation() { let validator = MergeValidator::new(); // Odd number of values assert_eq!(validator.compute_median(&[1.0, 3.0, 2.0]), 2.0); // Even number of values assert_eq!(validator.compute_median(&[1.0, 2.0, 3.0, 4.0]), 2.5); // Single value assert_eq!(validator.compute_median(&[5.0]), 5.0); // Empty values assert_eq!(validator.compute_median(&[]), 0.0); } #[test] fn test_majority_vote_resolution() -> Result<()> { let validator = MergeValidator::new(); // More positive values let result = validator.resolve_by_majority_vote(&[1.0, 2.0, -1.0])?; assert!(result > 0.0); // More negative values let result = validator.resolve_by_majority_vote(&[-1.0, -2.0, 1.0])?; assert!(result < 0.0); // Equal positive/negative (should average) let result = validator.resolve_by_majority_vote(&[1.0, -1.0])?; assert_eq!(result, 0.0); Ok(()) } }