Initial commit
This commit is contained in:
@@ -0,0 +1,540 @@
|
||||
//! 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(())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user