541 lines
18 KiB
Rust
541 lines
18 KiB
Rust
//! Advanced model merging strategies
|
|
//!
|
|
//! This module implements sophisticated merging strategies that build upon
|
|
//! the core algorithms:
|
|
//! - Model Soups: Ensemble averaging with greedy selection
|
|
//! - Frankenmerging: Layer-wise merging across architectures
|
|
//! - Progressive: Gradual merging with validation checkpoints
|
|
|
|
pub mod frankenmerge;
|
|
pub mod model_soup;
|
|
pub mod progressive;
|
|
|
|
// Re-export for convenience
|
|
pub use frankenmerge::*;
|
|
pub use model_soup::*;
|
|
pub use progressive::*;
|
|
|
|
use crate::error::Result;
|
|
use crate::types::{MergedModel, Model};
|
|
use std::collections::HashMap;
|
|
|
|
/// Common utilities for advanced merging strategies
|
|
pub struct StrategyUtils;
|
|
|
|
impl StrategyUtils {
|
|
/// Validate that all models in a group have compatible architectures
|
|
pub fn validate_model_group(models: &[Model]) -> Result<()> {
|
|
if models.len() < 2 {
|
|
return Ok(());
|
|
}
|
|
|
|
let reference = &models[0];
|
|
for (i, model) in models.iter().enumerate().skip(1) {
|
|
if !reference.is_compatible_with(model) {
|
|
return Err(crate::error::MergeError::compatibility(format!(
|
|
"Model {i} is incompatible with reference model in group"
|
|
)));
|
|
}
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Compute pairwise similarities between models
|
|
pub fn compute_pairwise_similarities(models: &[Model]) -> Result<Vec<Vec<f32>>> {
|
|
let n = models.len();
|
|
let mut similarities = vec![vec![0.0; n]; n];
|
|
|
|
for i in 0..n {
|
|
similarities[i][i] = 1.0; // Self-similarity
|
|
|
|
for j in (i + 1)..n {
|
|
let similarity =
|
|
crate::algorithms::MergeUtils::compute_similarity(&models[i], &models[j])?;
|
|
similarities[i][j] = similarity;
|
|
similarities[j][i] = similarity; // Symmetric
|
|
}
|
|
}
|
|
|
|
Ok(similarities)
|
|
}
|
|
|
|
/// Group models by similarity threshold
|
|
pub fn group_models_by_similarity(models: &[Model], threshold: f32) -> Result<Vec<Vec<usize>>> {
|
|
let similarities = Self::compute_pairwise_similarities(models)?;
|
|
let n = models.len();
|
|
|
|
let mut groups = Vec::new();
|
|
let mut assigned = vec![false; n];
|
|
|
|
for i in 0..n {
|
|
if assigned[i] {
|
|
continue;
|
|
}
|
|
|
|
let mut group = vec![i];
|
|
assigned[i] = true;
|
|
|
|
// Find all models similar to model i
|
|
for j in (i + 1)..n {
|
|
if !assigned[j] && similarities[i][j] >= threshold {
|
|
group.push(j);
|
|
assigned[j] = true;
|
|
}
|
|
}
|
|
|
|
groups.push(group);
|
|
}
|
|
|
|
Ok(groups)
|
|
}
|
|
|
|
/// Compute model diversity within a group
|
|
pub fn compute_group_diversity(models: &[Model], group_indices: &[usize]) -> Result<f32> {
|
|
if group_indices.len() < 2 {
|
|
return Ok(0.0);
|
|
}
|
|
|
|
let group_models: Vec<&Model> = group_indices.iter().map(|&i| &models[i]).collect();
|
|
let similarities = Self::compute_pairwise_similarities(
|
|
&group_models.into_iter().cloned().collect::<Vec<_>>(),
|
|
)?;
|
|
|
|
let mut total_dissimilarity = 0.0;
|
|
let mut pair_count = 0;
|
|
|
|
for i in 0..similarities.len() {
|
|
for j in (i + 1)..similarities.len() {
|
|
total_dissimilarity += 1.0 - similarities[i][j];
|
|
pair_count += 1;
|
|
}
|
|
}
|
|
|
|
Ok(if pair_count > 0 {
|
|
total_dissimilarity / pair_count as f32
|
|
} else {
|
|
0.0
|
|
})
|
|
}
|
|
|
|
/// Select representative models from a group
|
|
pub fn select_representatives(
|
|
models: &[Model],
|
|
group_indices: &[usize],
|
|
max_representatives: usize,
|
|
) -> Result<Vec<usize>> {
|
|
if group_indices.len() <= max_representatives {
|
|
return Ok(group_indices.to_vec());
|
|
}
|
|
|
|
// Use k-means-like selection based on parameter space
|
|
let mut representatives = Vec::new();
|
|
let mut remaining: Vec<usize> = group_indices.to_vec();
|
|
|
|
// First representative: model with highest average similarity to others
|
|
let mut best_idx = 0;
|
|
let mut best_avg_sim = 0.0;
|
|
|
|
for (i, &model_idx) in remaining.iter().enumerate() {
|
|
let mut total_sim = 0.0;
|
|
for &other_idx in &remaining {
|
|
if model_idx != other_idx {
|
|
total_sim += crate::algorithms::MergeUtils::compute_similarity(
|
|
&models[model_idx],
|
|
&models[other_idx],
|
|
)?;
|
|
}
|
|
}
|
|
|
|
let avg_sim = total_sim / (remaining.len() - 1) as f32;
|
|
if avg_sim > best_avg_sim {
|
|
best_avg_sim = avg_sim;
|
|
best_idx = i;
|
|
}
|
|
}
|
|
|
|
representatives.push(remaining.remove(best_idx));
|
|
|
|
// Subsequent representatives: maximize diversity
|
|
while representatives.len() < max_representatives && !remaining.is_empty() {
|
|
let mut best_idx = 0;
|
|
let mut best_min_sim = 1.0f32;
|
|
|
|
for (i, &candidate_idx) in remaining.iter().enumerate() {
|
|
let mut min_sim = 1.0f32;
|
|
|
|
for &rep_idx in &representatives {
|
|
let sim = crate::algorithms::MergeUtils::compute_similarity(
|
|
&models[candidate_idx],
|
|
&models[rep_idx],
|
|
)?;
|
|
min_sim = min_sim.min(sim);
|
|
}
|
|
|
|
if min_sim < best_min_sim {
|
|
best_min_sim = min_sim;
|
|
best_idx = i;
|
|
}
|
|
}
|
|
|
|
representatives.push(remaining.remove(best_idx));
|
|
}
|
|
|
|
Ok(representatives)
|
|
}
|
|
|
|
/// Compute model performance scores (placeholder - would use actual validation)
|
|
pub fn compute_performance_scores(models: &[Model]) -> HashMap<String, f32> {
|
|
let mut scores = HashMap::new();
|
|
|
|
for model in models {
|
|
// Placeholder scoring based on model complexity
|
|
let complexity = model.parameter_count() as f32;
|
|
let diversity = Self::compute_parameter_diversity(model);
|
|
|
|
// Simple heuristic score (in practice would use validation metrics)
|
|
let score = (1.0 / (1.0 + (complexity / 1000000.0))).max(0.1) * (0.5 + 0.5 * diversity);
|
|
|
|
scores.insert(model.name.clone(), score);
|
|
}
|
|
|
|
scores
|
|
}
|
|
|
|
/// Compute parameter diversity within a single model
|
|
pub fn compute_parameter_diversity(model: &Model) -> f32 {
|
|
let all_params: Vec<f32> = model
|
|
.parameters
|
|
.values()
|
|
.flat_map(|p| p.data.iter())
|
|
.copied()
|
|
.collect();
|
|
|
|
if all_params.is_empty() {
|
|
return 0.0;
|
|
}
|
|
|
|
let (_, std_dev, _, _) = crate::algorithms::MergeUtils::compute_moments(&all_params);
|
|
std_dev.min(1.0)
|
|
}
|
|
|
|
/// Create validation checkpoints for progressive merging
|
|
pub fn create_validation_checkpoints(
|
|
base_model: &Model,
|
|
target_model: &Model,
|
|
num_checkpoints: usize,
|
|
) -> Result<Vec<MergedModel>> {
|
|
let mut checkpoints = Vec::new();
|
|
|
|
for i in 0..=num_checkpoints {
|
|
let t = i as f32 / num_checkpoints as f32;
|
|
|
|
// Linear interpolation between base and target
|
|
let mut checkpoint_model = base_model.clone();
|
|
checkpoint_model.id = uuid::Uuid::new_v4();
|
|
checkpoint_model.name = format!("checkpoint_{i}_{t:.2}");
|
|
|
|
// Interpolate parameters
|
|
for (param_name, base_param) in &base_model.parameters {
|
|
if let Some(target_param) = target_model.parameters.get(param_name) {
|
|
let interpolated_data = crate::algorithms::MergeUtils::elementwise_op(
|
|
&base_param.data,
|
|
&target_param.data,
|
|
|base_val, target_val| (1.0 - t) * base_val + t * target_val,
|
|
)?;
|
|
|
|
if let Some(checkpoint_param) = checkpoint_model.parameters.get_mut(param_name)
|
|
{
|
|
checkpoint_param.data = interpolated_data;
|
|
}
|
|
}
|
|
}
|
|
|
|
// Create merged model wrapper
|
|
let merged_checkpoint = MergedModel {
|
|
model: checkpoint_model,
|
|
merge_info: crate::types::MergeInfo {
|
|
strategy: "ValidationCheckpoint".to_string(),
|
|
source_models: vec![
|
|
crate::types::ModelReference {
|
|
id: base_model.id,
|
|
name: base_model.name.clone(),
|
|
path: std::path::PathBuf::from(&base_model.name),
|
|
weight: Some(1.0 - t),
|
|
},
|
|
crate::types::ModelReference {
|
|
id: target_model.id,
|
|
name: target_model.name.clone(),
|
|
path: std::path::PathBuf::from(&target_model.name),
|
|
weight: Some(t),
|
|
},
|
|
],
|
|
merged_at: chrono::Utc::now(),
|
|
config: serde_json::json!({"t": t, "checkpoint": i}),
|
|
statistics: crate::types::MergeStatistics::default(),
|
|
},
|
|
quality_metrics: crate::types::QualityMetrics::default(),
|
|
};
|
|
|
|
checkpoints.push(merged_checkpoint);
|
|
}
|
|
|
|
Ok(checkpoints)
|
|
}
|
|
|
|
/// Estimate memory requirements for merging strategy
|
|
pub fn estimate_memory_requirements(models: &[Model], strategy: &str) -> Result<usize> {
|
|
let total_model_memory: usize = models.iter().map(super::types::Model::memory_size).sum();
|
|
|
|
let overhead_factor = match strategy {
|
|
"ModelSoup" => 1.2, // Small overhead for ensemble averaging
|
|
"Frankenmerge" => 1.5, // Moderate overhead for layer-wise operations
|
|
"Progressive" => 2.0, // Higher overhead for checkpoints
|
|
_ => 1.1,
|
|
};
|
|
|
|
let estimated_memory = (total_model_memory as f32 * overhead_factor) as usize;
|
|
Ok(estimated_memory)
|
|
}
|
|
|
|
/// Check resource constraints
|
|
pub fn check_resource_constraints(
|
|
models: &[Model],
|
|
strategy: &str,
|
|
max_memory_mb: Option<usize>,
|
|
) -> Result<()> {
|
|
let estimated_memory = Self::estimate_memory_requirements(models, strategy)?;
|
|
let estimated_memory_mb = estimated_memory / (1024 * 1024);
|
|
|
|
if let Some(max_mb) = max_memory_mb
|
|
&& estimated_memory_mb > max_mb
|
|
{
|
|
return Err(crate::error::MergeError::memory(estimated_memory_mb));
|
|
}
|
|
|
|
// Check for extremely large memory requirements
|
|
if estimated_memory_mb > 32768 {
|
|
// 32GB
|
|
return Err(crate::error::MergeError::resource_exhausted(format!(
|
|
"Memory requirement: {estimated_memory_mb}MB"
|
|
)));
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
#[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: std::collections::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
|
|
}
|
|
|
|
#[test]
|
|
fn test_model_group_validation() -> Result<()> {
|
|
let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0])]);
|
|
|
|
let model2 = create_test_model("model2", vec![("weight".to_string(), vec![3.0, 4.0])]);
|
|
|
|
let result = StrategyUtils::validate_model_group(&[model1, model2]);
|
|
assert!(result.is_ok());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_pairwise_similarities() -> 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.0, 2.0, 3.0]), // Identical
|
|
],
|
|
);
|
|
|
|
let model3 = create_test_model(
|
|
"model3",
|
|
vec![
|
|
("weight".to_string(), vec![4.0, 5.0, 6.0]), // Different
|
|
],
|
|
);
|
|
|
|
let similarities = StrategyUtils::compute_pairwise_similarities(&[model1, model2, model3])?;
|
|
|
|
assert_eq!(similarities.len(), 3);
|
|
assert_eq!(similarities[0].len(), 3);
|
|
|
|
// Self-similarities should be 1.0
|
|
assert_eq!(similarities[0][0], 1.0);
|
|
assert_eq!(similarities[1][1], 1.0);
|
|
assert_eq!(similarities[2][2], 1.0);
|
|
|
|
// Model1 and Model2 should be very similar (identical)
|
|
assert!(similarities[0][1] > 0.9);
|
|
|
|
// Model1/Model2 and Model3 should be less similar
|
|
assert!(similarities[0][2] < similarities[0][1]);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_similarity_grouping() -> Result<()> {
|
|
let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0, 2.0])]);
|
|
|
|
let model2 = create_test_model(
|
|
"model2",
|
|
vec![
|
|
("weight".to_string(), vec![1.1, 2.1]), // Very similar to model1
|
|
],
|
|
);
|
|
|
|
let model3 = create_test_model(
|
|
"model3",
|
|
vec![
|
|
("weight".to_string(), vec![10.0, 20.0]), // Very different
|
|
],
|
|
);
|
|
|
|
let groups = StrategyUtils::group_models_by_similarity(&[model1, model2, model3], 0.8)?;
|
|
|
|
// Should have at least 2 groups (similar models together, different model separate)
|
|
assert!(groups.len() >= 2);
|
|
|
|
// Total models should be preserved
|
|
let total_models: usize = groups.iter().map(|g| g.len()).sum();
|
|
assert_eq!(total_models, 3);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_representative_selection() -> Result<()> {
|
|
let models = (0..5)
|
|
.map(|i| {
|
|
create_test_model(
|
|
&format!("model{}", i),
|
|
vec![("weight".to_string(), vec![i as f32, (i + 1) as f32])],
|
|
)
|
|
})
|
|
.collect::<Vec<_>>();
|
|
|
|
let group_indices = vec![0, 1, 2, 3, 4];
|
|
let representatives = StrategyUtils::select_representatives(&models, &group_indices, 3)?;
|
|
|
|
assert_eq!(representatives.len(), 3);
|
|
assert!(representatives.iter().all(|&i| i < 5));
|
|
|
|
// Should select diverse representatives
|
|
assert!(representatives[0] != representatives[1]);
|
|
assert!(representatives[1] != representatives[2]);
|
|
assert!(representatives[0] != representatives[2]);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_parameter_diversity() {
|
|
let model1 = create_test_model(
|
|
"uniform",
|
|
vec![
|
|
("weight".to_string(), vec![1.0, 1.0, 1.0, 1.0]), // No diversity
|
|
],
|
|
);
|
|
|
|
let model2 = create_test_model(
|
|
"diverse",
|
|
vec![
|
|
("weight".to_string(), vec![1.0, 5.0, 10.0, 2.0]), // High diversity
|
|
],
|
|
);
|
|
|
|
let diversity1 = StrategyUtils::compute_parameter_diversity(&model1);
|
|
let diversity2 = StrategyUtils::compute_parameter_diversity(&model2);
|
|
|
|
assert!(diversity1 < diversity2);
|
|
assert_eq!(diversity1, 0.0); // Uniform values should have zero diversity
|
|
assert!(diversity2 > 0.0); // Diverse values should have positive diversity
|
|
}
|
|
|
|
#[test]
|
|
fn test_memory_estimation() -> Result<()> {
|
|
let model1 = create_test_model("model1", vec![("weight".to_string(), vec![1.0; 1000])]);
|
|
|
|
let model2 = create_test_model("model2", vec![("weight".to_string(), vec![1.0; 1000])]);
|
|
|
|
let models = vec![model1, model2];
|
|
|
|
let soup_memory = StrategyUtils::estimate_memory_requirements(&models, "ModelSoup")?;
|
|
let progressive_memory =
|
|
StrategyUtils::estimate_memory_requirements(&models, "Progressive")?;
|
|
|
|
// Progressive should require more memory than ModelSoup
|
|
assert!(progressive_memory > soup_memory);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_resource_constraints() -> Result<()> {
|
|
let small_model = create_test_model("small", vec![("weight".to_string(), vec![1.0; 10])]);
|
|
|
|
let models = vec![small_model];
|
|
|
|
// Should pass with reasonable memory limit
|
|
let result = StrategyUtils::check_resource_constraints(&models, "ModelSoup", Some(1024));
|
|
assert!(result.is_ok());
|
|
|
|
// Should fail with very small memory limit
|
|
let result = StrategyUtils::check_resource_constraints(&models, "ModelSoup", Some(1));
|
|
assert!(result.is_err());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_validation_checkpoints() -> Result<()> {
|
|
let base_model = create_test_model("base", vec![("weight".to_string(), vec![0.0, 0.0])]);
|
|
|
|
let target_model =
|
|
create_test_model("target", vec![("weight".to_string(), vec![2.0, 4.0])]);
|
|
|
|
let checkpoints =
|
|
StrategyUtils::create_validation_checkpoints(&base_model, &target_model, 4)?;
|
|
|
|
assert_eq!(checkpoints.len(), 5); // 0, 1, 2, 3, 4 (inclusive)
|
|
|
|
// First checkpoint should be base model
|
|
let first_weight = checkpoints[0].model.get_parameter("weight").unwrap();
|
|
assert_eq!(first_weight.data, vec![0.0, 0.0]);
|
|
|
|
// Last checkpoint should be target model
|
|
let last_weight = checkpoints[4].model.get_parameter("weight").unwrap();
|
|
assert_eq!(last_weight.data, vec![2.0, 4.0]);
|
|
|
|
// Middle checkpoint should be interpolated
|
|
let middle_weight = checkpoints[2].model.get_parameter("weight").unwrap();
|
|
assert_eq!(middle_weight.data, vec![1.0, 2.0]); // 50% interpolation
|
|
|
|
Ok(())
|
|
}
|
|
}
|