Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -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 &param.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(&param.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", &param_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", &param_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", &param_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(())
}
}