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,592 @@
//! TIES (Task-specific Interference Elimination) algorithm implementation
//!
//! TIES resolves parameter interference by:
//! 1. Eliminating redundant parameters through sign voting
//! 2. Selecting top-k parameters by magnitude
//! 3. Rescaling merged parameters appropriately
//!
//! Reference: "Resolving Interference When Merging Models" (Yadav et al., 2024)
use crate::algorithms::MergeUtils;
use crate::config::{RescaleMethod, TiesConfig};
use crate::error::{MergeError, Result};
use crate::types::{
MergeInfo, MergeStatistics, MergedModel, Model, ModelReference, QualityMetrics,
};
use indexmap::IndexMap;
use std::collections::HashMap;
use tracing::{debug, info};
/// TIES merging algorithm implementation
pub struct TiesMerger {
config: TiesConfig,
}
impl TiesMerger {
/// Create a new TIES merger with configuration
pub fn new(config: TiesConfig) -> Self {
Self { config }
}
/// Execute TIES merging on multiple models
pub async fn merge(&self, models: &[Model]) -> Result<MergedModel> {
if models.len() < 2 {
return Err(MergeError::algorithm(
"TIES",
"At least 2 models required for merging",
));
}
info!("Starting TIES merge of {} models", models.len());
let start_time = std::time::Instant::now();
// Step 1: Validate model compatibility
self.validate_compatibility(models)?;
// Step 2: Compute task vectors (differences from base model)
let task_vectors = self.compute_task_vectors(models)?;
// Step 3: Apply TIES algorithm
let merged_parameters = self.apply_ties_algorithm(&task_vectors, models)?;
// Step 4: Create merged model
let mut merged_model = models[0].clone();
merged_model.id = uuid::Uuid::new_v4();
merged_model.name = format!("TIES_merged_{}", models.len());
merged_model.parameters = merged_parameters;
// Step 5: Compute merge statistics and quality metrics
let statistics = self.compute_statistics(&merged_model, models, start_time.elapsed());
let quality_metrics = self.compute_quality_metrics(&merged_model, models)?;
let merge_info = MergeInfo {
strategy: "TIES".to_string(),
source_models: models
.iter()
.map(|m| ModelReference {
id: m.id,
name: m.name.clone(),
path: std::path::PathBuf::from(&m.name),
weight: Some(1.0 / models.len() as f32),
})
.collect(),
merged_at: chrono::Utc::now(),
config: serde_json::to_value(&self.config)
.map_err(|e| MergeError::internal(format!("Config serialization: {e}")))?,
statistics,
};
info!("TIES merge completed in {:?}", start_time.elapsed());
Ok(MergedModel {
model: merged_model,
merge_info,
quality_metrics,
})
}
/// Validate that models are compatible for TIES merging
fn validate_compatibility(&self, models: &[Model]) -> Result<()> {
let base_model = &models[0];
for (i, model) in models.iter().enumerate().skip(1) {
if !base_model.is_compatible_with(model) {
return Err(MergeError::compatibility(format!(
"Model {i} incompatible with base model"
)));
}
}
Ok(())
}
/// Compute task vectors (parameter differences from base model)
fn compute_task_vectors(&self, models: &[Model]) -> Result<Vec<IndexMap<String, Vec<f32>>>> {
info!("Computing task vectors from {} models", models.len());
let base_model = &models[0];
let mut task_vectors = Vec::new();
for model in models.iter().skip(1) {
let mut task_vector = IndexMap::new();
for (param_name, base_param) in &base_model.parameters {
if let Some(model_param) = model.parameters.get(param_name) {
let diff =
MergeUtils::elementwise_op(&model_param.data, &base_param.data, |a, b| {
a - b
})?;
task_vector.insert(param_name.clone(), diff);
}
}
task_vectors.push(task_vector);
}
Ok(task_vectors)
}
/// Apply the TIES algorithm to merge task vectors
fn apply_ties_algorithm(
&self,
task_vectors: &[IndexMap<String, Vec<f32>>],
models: &[Model],
) -> Result<IndexMap<String, crate::types::ParameterTensor>> {
info!(
"Applying TIES algorithm with density: {}",
self.config.density
);
let base_model = &models[0];
let mut merged_parameters = IndexMap::new();
let mut total_conflicts = 0;
let mut total_dropped = 0;
for (param_name, base_param) in &base_model.parameters {
debug!("Processing parameter: {}", param_name);
// Collect all task vectors for this parameter
let param_vectors: Vec<&Vec<f32>> = task_vectors
.iter()
.filter_map(|tv| tv.get(param_name))
.collect();
if param_vectors.is_empty() {
// Keep base parameter if no task vectors available
merged_parameters.insert(param_name.clone(), base_param.clone());
continue;
}
// Step 1: Sign consistency check
let (consistent_values, conflicts) = if self.config.enable_sign_voting {
self.resolve_sign_conflicts(&param_vectors, param_name)?
} else {
// Simple average without sign checking
let param_slices: Vec<&[f32]> =
param_vectors.iter().map(|v| v.as_slice()).collect();
let avg = MergeUtils::average_parameters(&param_slices)?;
(avg, 0)
};
total_conflicts += conflicts;
// Step 2: Magnitude-based parameter selection
let importance_scores = MergeUtils::compute_magnitude_importance(&consistent_values);
let k = (consistent_values.len() as f32 * self.config.density) as usize;
let selection_mask = MergeUtils::select_top_k_parameters(&importance_scores, k);
// Count dropped parameters
let dropped = selection_mask.iter().filter(|&&x| !x).count();
total_dropped += dropped;
// Step 3: Apply selection mask and rescaling
let mut final_values = consistent_values.clone();
MergeUtils::apply_sparsity_mask(&mut final_values, &selection_mask)?;
// Rescale if needed
match self.config.rescale_method {
RescaleMethod::Magnitude => {
self.rescale_by_magnitude(&mut final_values, &selection_mask)?;
}
RescaleMethod::SignConsistency => {
self.rescale_by_sign_consistency(&mut final_values, &param_vectors)?;
}
RescaleMethod::None => {
// No rescaling
}
}
// Add back to base parameters
let merged_values =
MergeUtils::elementwise_op(&base_param.data, &final_values, |base, delta| {
base + delta
})?;
// Create merged parameter tensor
let mut merged_param = base_param.clone();
merged_param.data = merged_values;
merged_parameters.insert(param_name.clone(), merged_param);
}
info!(
"TIES algorithm completed. Conflicts resolved: {}, Parameters dropped: {}",
total_conflicts, total_dropped
);
Ok(merged_parameters)
}
/// Resolve sign conflicts using voting mechanism
fn resolve_sign_conflicts(
&self,
param_vectors: &[&Vec<f32>],
param_name: &str,
) -> Result<(Vec<f32>, usize)> {
if param_vectors.is_empty() {
return Ok((vec![], 0));
}
let param_len = param_vectors[0].len();
let mut resolved_values = vec![0.0; param_len];
let mut conflicts = 0;
debug!(
"Resolving sign conflicts for parameter: {} ({} vectors, {} elements)",
param_name,
param_vectors.len(),
param_len
);
// Process each parameter element
for i in 0..param_len {
let values: Vec<f32> = param_vectors.iter().map(|v| v[i]).collect();
// Check sign consistency
let positive_count = values.iter().filter(|&&x| x > 0.0).count();
let negative_count = values.iter().filter(|&&x| x < 0.0).count();
let _zero_count = values.iter().filter(|&&x| x == 0.0).count();
let total_nonzero = positive_count + negative_count;
if total_nonzero == 0 {
resolved_values[i] = 0.0;
continue;
}
// Determine consensus sign
let consensus_positive =
positive_count as f32 / total_nonzero as f32 >= self.config.voting_threshold;
let consensus_negative =
negative_count as f32 / total_nonzero as f32 >= self.config.voting_threshold;
if consensus_positive && !consensus_negative {
// Use only positive values
let positive_values: Vec<f32> =
values.iter().filter(|&&x| x > 0.0).copied().collect();
if !positive_values.is_empty() {
resolved_values[i] =
positive_values.iter().sum::<f32>() / positive_values.len() as f32;
}
if negative_count > 0 {
conflicts += negative_count;
}
} else if consensus_negative && !consensus_positive {
// Use only negative values
let negative_values: Vec<f32> =
values.iter().filter(|&&x| x < 0.0).copied().collect();
if !negative_values.is_empty() {
resolved_values[i] =
negative_values.iter().sum::<f32>() / negative_values.len() as f32;
}
if positive_count > 0 {
conflicts += positive_count;
}
} else {
// No clear consensus - use magnitude-weighted average
let total_magnitude: f32 = values.iter().map(|x| x.abs()).sum();
if total_magnitude > 0.0 {
resolved_values[i] =
values.iter().map(|&x| x * x.abs()).sum::<f32>() / total_magnitude;
}
conflicts += std::cmp::min(positive_count, negative_count);
}
}
debug!(
"Sign conflict resolution completed for {}: {} conflicts resolved",
param_name, conflicts
);
Ok((resolved_values, conflicts))
}
/// Rescale parameters based on magnitude
fn rescale_by_magnitude(&self, values: &mut [f32], mask: &[bool]) -> Result<()> {
let active_count = mask.iter().filter(|&&x| x).count();
if active_count == 0 {
return Ok(());
}
let total_count = mask.len();
let rescale_factor = (total_count as f32) / (active_count as f32);
for (value, &active) in values.iter_mut().zip(mask.iter()) {
if active {
*value *= rescale_factor;
}
}
Ok(())
}
/// Rescale parameters based on sign consistency
fn rescale_by_sign_consistency(
&self,
values: &mut [f32],
param_vectors: &[&Vec<f32>],
) -> Result<()> {
if param_vectors.is_empty() {
return Ok(());
}
let param_slices: Vec<&[f32]> = param_vectors.iter().map(|v| v.as_slice()).collect();
let consistency_scores =
MergeUtils::check_sign_consistency(&param_slices, self.config.sign_threshold);
let consistent_count = consistency_scores.iter().filter(|&&x| x).count();
if consistent_count == 0 {
return Ok(());
}
let rescale_factor = (consistency_scores.len() as f32) / (consistent_count as f32);
for (value, &consistent) in values.iter_mut().zip(consistency_scores.iter()) {
if consistent {
*value *= rescale_factor;
}
}
Ok(())
}
/// Compute merge statistics
fn compute_statistics(
&self,
merged_model: &Model,
_source_models: &[Model],
duration: std::time::Duration,
) -> MergeStatistics {
let parameters_merged = merged_model.parameter_count();
let memory_usage_mb = merged_model.memory_size() / (1024 * 1024);
MergeStatistics {
parameters_merged,
parameters_conflicted: 0, // Would be tracked during merging
parameters_dropped: 0, // Would be tracked during merging
duration_ms: duration.as_millis() as u64,
memory_usage_mb,
gpu_memory_mb: None,
}
}
/// Compute quality metrics for the merged model
fn compute_quality_metrics(
&self,
merged_model: &Model,
source_models: &[Model],
) -> Result<QualityMetrics> {
let mut consistency_scores = Vec::new();
// Compute consistency with each source model
for source in source_models {
let similarity = MergeUtils::compute_similarity(merged_model, source)?;
consistency_scores.push(similarity);
}
let consistency_score =
consistency_scores.iter().sum::<f32>() / consistency_scores.len() as f32;
// Compute complexity score (normalized parameter variance)
let all_params: Vec<f32> = merged_model
.parameters
.values()
.flat_map(|p| p.data.iter())
.copied()
.collect();
let (_, std_dev, _, _) = MergeUtils::compute_moments(&all_params);
let complexity_score = std_dev.min(1.0); // Normalized to [0, 1]
let mut quality_indicators = HashMap::new();
quality_indicators.insert("parameter_diversity".to_string(), std_dev);
quality_indicators.insert("average_similarity".to_string(), consistency_score);
let mut validation_results = HashMap::new();
validation_results.insert("sign_consistency".to_string(), consistency_score > 0.5);
validation_results.insert("magnitude_preservation".to_string(), complexity_score > 0.1);
let mut performance_predictions = HashMap::new();
performance_predictions.insert("expected_accuracy".to_string(), consistency_score * 0.9);
performance_predictions.insert("stability_score".to_string(), 1.0 - complexity_score);
Ok(QualityMetrics {
consistency_score,
complexity_score,
quality_indicators,
validation_results,
performance_predictions,
})
}
}
/// Convenience function for TIES merging
pub async fn merge_models(models: &[Model], config: &TiesConfig) -> Result<MergedModel> {
let merger = TiesMerger::new(config.clone());
merger.merge(models).await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{DataType, ModelArchitecture, ParameterTensor};
fn create_test_model(name: &str, params: Vec<(String, Vec<f32>)>) -> Model {
let arch = ModelArchitecture {
arch_type: "test".to_string(),
num_layers: 1,
hidden_dim: params.len(),
params: HashMap::new(),
};
let mut model = Model::new(name.to_string(), arch);
for (param_name, data) in params {
let param = ParameterTensor::new(param_name, vec![data.len()], DataType::Float32, data);
model.add_parameter(param);
}
model
}
#[tokio::test]
async fn test_ties_merge_basic() -> Result<()> {
let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0, 3.0])]);
let model2 = create_test_model("model2", vec![("weight".to_string(), vec![1.1, 2.1, 3.1])]);
let config = TiesConfig::default();
let merger = TiesMerger::new(config);
let result = merger.merge(&[model1, model2]).await?;
assert!(result.model.parameters.contains_key("weight"));
assert_eq!(result.merge_info.strategy, "TIES");
Ok(())
}
#[tokio::test]
async fn test_ties_sign_conflict_resolution() -> Result<()> {
let merger = TiesMerger::new(TiesConfig {
enable_sign_voting: true,
voting_threshold: 0.6,
..TiesConfig::default()
});
let vec1 = vec![1.0, -2.0, 3.0];
let vec2 = vec![2.0, -1.0, 4.0];
let vec3 = vec![-1.0, -3.0, -2.0];
let param_vectors = vec![&vec1, &vec2, &vec3];
let (resolved, conflicts) = merger.resolve_sign_conflicts(&param_vectors, "test")?;
assert!(conflicts > 0);
assert_eq!(resolved.len(), 3);
// Second element should be consistently negative
assert!(resolved[1] < 0.0);
Ok(())
}
#[tokio::test]
async fn test_ties_density_selection() -> Result<()> {
let model1 = create_test_model(
"model1",
vec![("weight".to_string(), vec![0.0, 0.0, 0.0, 0.0, 0.0])],
);
let model2 = create_test_model(
"model2",
vec![
("weight".to_string(), vec![1.0, 0.1, 0.5, 0.05, 0.8]), // Different magnitudes
],
);
let config = TiesConfig {
density: 0.6, // Keep top 60% = 3 parameters
..TiesConfig::default()
};
let merger = TiesMerger::new(config);
let result = merger.merge(&[model1, model2]).await?;
let merged_weight = result.model.get_parameter("weight").unwrap();
let non_zero_count = merged_weight.data.iter().filter(|&&x| x != 0.0).count();
// Should have at most 3 non-zero parameters (top 60%)
assert!(non_zero_count <= 3);
Ok(())
}
#[tokio::test]
async fn test_ties_rescaling() -> Result<()> {
let model1 = create_test_model("model1", vec![("weight".to_string(), vec![0.0, 0.0])]);
let model2 = create_test_model("model2", vec![("weight".to_string(), vec![1.0, 1.0])]);
let config = TiesConfig {
density: 0.5, // Keep only 50% of parameters
rescale_method: RescaleMethod::Magnitude,
..TiesConfig::default()
};
let merger = TiesMerger::new(config);
let result = merger.merge(&[model1, model2]).await?;
let merged_weight = result.model.get_parameter("weight").unwrap();
// Check that rescaling was applied
let non_zero_values: Vec<f32> = merged_weight
.data
.iter()
.filter(|&&x| x != 0.0)
.cloned()
.collect();
if !non_zero_values.is_empty() {
// Rescaled values should be larger than original due to density < 1.0
assert!(non_zero_values.iter().any(|&x| x > 1.0));
}
Ok(())
}
#[test]
fn test_task_vector_computation() -> Result<()> {
let model1 = create_test_model("base", vec![("weight".to_string(), vec![1.0, 2.0, 3.0])]);
let model2 = create_test_model("task", vec![("weight".to_string(), vec![2.0, 3.0, 4.0])]);
let merger = TiesMerger::new(TiesConfig::default());
let task_vectors = merger.compute_task_vectors(&[model1, model2])?;
assert_eq!(task_vectors.len(), 1);
assert!(task_vectors[0].contains_key("weight"));
let task_vector = &task_vectors[0]["weight"];
assert_eq!(task_vector, &vec![1.0, 1.0, 1.0]);
Ok(())
}
#[test]
fn test_quality_metrics_computation() -> Result<()> {
let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0, 3.0])]);
let model2 = create_test_model("model2", vec![("weight".to_string(), vec![1.1, 2.1, 3.1])]);
let merged = create_test_model(
"merged",
vec![("weight".to_string(), vec![1.05, 2.05, 3.05])],
);
let merger = TiesMerger::new(TiesConfig::default());
let metrics = merger.compute_quality_metrics(&merged, &[model1, model2])?;
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());
Ok(())
}
}