Initial commit
This commit is contained in:
@@ -0,0 +1,452 @@
|
||||
//! 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<Vec<ParameterConflict>> {
|
||||
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<Vec<f32>> {
|
||||
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<f32> = 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::<f32>() / 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<Vec<ParameterConflict>> {
|
||||
let mut conflicts = Vec::new();
|
||||
let param_len = param_vectors[0].len();
|
||||
|
||||
for i in 0..param_len {
|
||||
let values: Vec<f32> = 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<Vec<ParameterConflict>> {
|
||||
let mut conflicts = Vec::new();
|
||||
let param_len = param_vectors[0].len();
|
||||
|
||||
for i in 0..param_len {
|
||||
let values: Vec<f32> = 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<Vec<ParameterConflict>> {
|
||||
let mut conflicts = Vec::new();
|
||||
let param_len = param_vectors[0].len();
|
||||
|
||||
for i in 0..param_len {
|
||||
let values: Vec<f32> = 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<f32> {
|
||||
if conflicts.is_empty() {
|
||||
return Ok(values.iter().sum::<f32>() / 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::<f32>() / values.len() as f32),
|
||||
ConflictResolution::WeightedAverage => {
|
||||
// Use inverse severity as weights
|
||||
let weights: Vec<f32> = 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::<f32>() / values.len() as f32),
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_by_majority_vote(&self, values: &[f32]) -> Result<f32> {
|
||||
if values.is_empty() {
|
||||
return Ok(0.0);
|
||||
}
|
||||
|
||||
// Sign-based majority vote
|
||||
let positive_values: Vec<f32> = values.iter().filter(|&&x| x > 0.0).copied().collect();
|
||||
let negative_values: Vec<f32> = values.iter().filter(|&&x| x < 0.0).copied().collect();
|
||||
|
||||
if positive_values.len() > negative_values.len() {
|
||||
Ok(positive_values.iter().sum::<f32>() / positive_values.len() as f32)
|
||||
} else if negative_values.len() > positive_values.len() {
|
||||
Ok(negative_values.iter().sum::<f32>() / negative_values.len() as f32)
|
||||
} else {
|
||||
// Tie - use simple average
|
||||
Ok(values.iter().sum::<f32>() / 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(())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user