Initial commit
This commit is contained in:
@@ -0,0 +1,240 @@
|
||||
//! Frankenmerge: Layer-wise model merging across architectures
|
||||
//!
|
||||
//! This strategy allows combining layers from different models,
|
||||
//! even with slight architectural differences.
|
||||
|
||||
use crate::config::FrankenmergeConfig;
|
||||
use crate::error::{MergeError, Result};
|
||||
use crate::types::{
|
||||
MergeInfo, MergeStatistics, MergedModel, Model, ModelReference, ParameterTensor, QualityMetrics,
|
||||
};
|
||||
use indexmap::IndexMap;
|
||||
use std::collections::HashMap;
|
||||
use uuid::Uuid;
|
||||
|
||||
/// Frankenmerge implementation
|
||||
pub struct Frankenmerge {
|
||||
config: FrankenmergeConfig,
|
||||
layer_compatibility_cache: HashMap<(String, String), f32>,
|
||||
}
|
||||
|
||||
impl Frankenmerge {
|
||||
/// Create a new Frankenmerge instance
|
||||
pub fn new(config: FrankenmergeConfig) -> Self {
|
||||
Self {
|
||||
config,
|
||||
layer_compatibility_cache: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge models using Frankenmerge strategy
|
||||
pub fn merge(&mut self, models: Vec<Model>) -> Result<MergedModel> {
|
||||
if models.len() < 2 {
|
||||
return Err(MergeError::config(
|
||||
"models",
|
||||
"Frankenmerge requires at least 2 models",
|
||||
));
|
||||
}
|
||||
|
||||
let mut merged_parameters = IndexMap::new();
|
||||
let base_model = &models[0];
|
||||
|
||||
// Process each layer in the base model
|
||||
for (layer_name, base_param) in &base_model.parameters {
|
||||
let compatible_layers =
|
||||
self.find_compatible_layers(layer_name, base_param, &models[1..])?;
|
||||
|
||||
if !compatible_layers.is_empty() {
|
||||
let merged_param = self.merge_compatible_layers(base_param, compatible_layers)?;
|
||||
merged_parameters.insert(layer_name.clone(), merged_param);
|
||||
} else {
|
||||
// Keep original layer if no compatible layers found
|
||||
merged_parameters.insert(layer_name.clone(), base_param.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Add unique layers from other models
|
||||
for model in &models[1..] {
|
||||
for (layer_name, param) in &model.parameters {
|
||||
if !merged_parameters.contains_key(layer_name) {
|
||||
merged_parameters.insert(layer_name.clone(), param.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let merged_model = Model {
|
||||
id: Uuid::new_v4(),
|
||||
name: format!("frankenmerged_{}", chrono::Utc::now().timestamp()),
|
||||
architecture: base_model.architecture.clone(),
|
||||
parameters: merged_parameters,
|
||||
metadata: base_model.metadata.clone(),
|
||||
config: base_model.config.clone(),
|
||||
};
|
||||
|
||||
let source_refs: Vec<ModelReference> = models
|
||||
.iter()
|
||||
.map(|m| ModelReference {
|
||||
id: m.id,
|
||||
name: m.name.clone(),
|
||||
path: std::path::PathBuf::from(&m.name),
|
||||
weight: None,
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(MergedModel {
|
||||
model: merged_model,
|
||||
merge_info: MergeInfo {
|
||||
strategy: "frankenmerge".to_string(),
|
||||
source_models: source_refs,
|
||||
merged_at: chrono::Utc::now(),
|
||||
config: serde_json::to_value(&self.config).unwrap_or(serde_json::Value::Null),
|
||||
statistics: MergeStatistics {
|
||||
parameters_merged: base_model.parameters.len(),
|
||||
parameters_conflicted: 0,
|
||||
parameters_dropped: 0,
|
||||
duration_ms: 0,
|
||||
memory_usage_mb: 0,
|
||||
gpu_memory_mb: Some(0),
|
||||
},
|
||||
},
|
||||
quality_metrics: QualityMetrics {
|
||||
consistency_score: 1.0,
|
||||
complexity_score: 1.0,
|
||||
quality_indicators: std::collections::HashMap::new(),
|
||||
validation_results: std::collections::HashMap::new(),
|
||||
performance_predictions: std::collections::HashMap::new(),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// Find compatible layers across models
|
||||
fn find_compatible_layers<'a>(
|
||||
&self,
|
||||
target_layer: &'a str,
|
||||
target_param: &ParameterTensor,
|
||||
other_models: &'a [Model],
|
||||
) -> Result<Vec<(&'a str, &'a ParameterTensor)>> {
|
||||
let mut compatible = Vec::new();
|
||||
|
||||
for model in other_models {
|
||||
// Check same-named layer and use layer weights if available
|
||||
if let Some(param) = model.parameters.get(target_layer) {
|
||||
let compatibility = self.compute_compatibility(target_param, param)?;
|
||||
if compatibility >= 0.9 {
|
||||
// Default threshold
|
||||
compatible.push((target_layer, param));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(compatible)
|
||||
}
|
||||
|
||||
/// Compute compatibility score between two parameter tensors
|
||||
fn compute_compatibility(
|
||||
&self,
|
||||
param1: &ParameterTensor,
|
||||
param2: &ParameterTensor,
|
||||
) -> Result<f32> {
|
||||
// Skip cache for now to avoid mutability issues
|
||||
let score = if param1.shape == param2.shape {
|
||||
1.0 // Perfect compatibility
|
||||
} else if param1.shape.len() != param2.shape.len() {
|
||||
0.0 // Incompatible rank
|
||||
} else {
|
||||
// Compute dimensional compatibility
|
||||
let mut total_ratio = 1.0f32;
|
||||
for (d1, d2) in param1.shape.iter().zip(¶m2.shape) {
|
||||
let ratio = (*d1 as f32) / (*d2 as f32);
|
||||
total_ratio *= ratio.max(1.0 / ratio);
|
||||
}
|
||||
|
||||
if total_ratio <= 1.5 {
|
||||
// Default max dimension ratio
|
||||
1.0 / total_ratio // Higher score for closer dimensions
|
||||
} else {
|
||||
0.0
|
||||
}
|
||||
};
|
||||
|
||||
Ok(score)
|
||||
}
|
||||
|
||||
/// Merge compatible layers into a single parameter
|
||||
fn merge_compatible_layers(
|
||||
&self,
|
||||
base_param: &ParameterTensor,
|
||||
compatible_layers: Vec<(&str, &ParameterTensor)>,
|
||||
) -> Result<ParameterTensor> {
|
||||
if compatible_layers.is_empty() {
|
||||
return Ok(base_param.clone());
|
||||
}
|
||||
|
||||
// For now, simple averaging
|
||||
// In a full implementation, this would handle dimension interpolation
|
||||
let mut merged = base_param.clone();
|
||||
let total_weight = 1.0 + compatible_layers.len() as f32;
|
||||
|
||||
// Scale base parameter
|
||||
for value in &mut merged.data {
|
||||
*value /= total_weight;
|
||||
}
|
||||
|
||||
// Add contributions from compatible layers
|
||||
for (_, param) in compatible_layers {
|
||||
if param.shape == base_param.shape {
|
||||
for (i, value) in param.data.iter().enumerate() {
|
||||
if i < merged.data.len() {
|
||||
merged.data[i] += value / total_weight;
|
||||
}
|
||||
}
|
||||
}
|
||||
// Interpolation for different shapes would be implemented here
|
||||
}
|
||||
|
||||
Ok(merged)
|
||||
}
|
||||
}
|
||||
|
||||
/// Convenience function for Frankenmerge merging
|
||||
pub async fn merge_models(models: &[Model], config: &FrankenmergeConfig) -> Result<MergedModel> {
|
||||
let mut merger = Frankenmerge::new(config.clone());
|
||||
merger.merge(models.to_vec())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use uuid::Uuid;
|
||||
|
||||
#[test]
|
||||
fn test_frankenmerge_creation() {
|
||||
let config = FrankenmergeConfig::default();
|
||||
let frankenmerge = Frankenmerge::new(config);
|
||||
assert!(frankenmerge.layer_compatibility_cache.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_compatibility_computation() {
|
||||
let mut frankenmerge = Frankenmerge::new(FrankenmergeConfig::default());
|
||||
|
||||
let param1 = ParameterTensor::new(
|
||||
"param1".to_string(),
|
||||
vec![10, 20],
|
||||
crate::types::DataType::Float32,
|
||||
vec![0.0; 200],
|
||||
);
|
||||
|
||||
let param2 = ParameterTensor::new(
|
||||
"param2".to_string(),
|
||||
vec![10, 20],
|
||||
crate::types::DataType::Float32,
|
||||
vec![0.0; 200],
|
||||
);
|
||||
|
||||
let compatibility = frankenmerge
|
||||
.compute_compatibility(¶m1, ¶m2)
|
||||
.unwrap();
|
||||
assert_eq!(compatibility, 1.0);
|
||||
}
|
||||
}
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,361 @@
|
||||
//! Progressive merging: Gradual model merging with validation checkpoints
|
||||
//!
|
||||
//! This strategy gradually merges models over multiple steps, allowing
|
||||
//! for validation and adjustment at each checkpoint.
|
||||
|
||||
use crate::config::ProgressiveConfig;
|
||||
use crate::error::{MergeError, Result};
|
||||
use crate::types::{
|
||||
MergeInfo, MergeStatistics, MergedModel, Model, ParameterTensor, QualityMetrics,
|
||||
};
|
||||
use indexmap::IndexMap;
|
||||
use std::sync::Arc;
|
||||
use uuid::Uuid;
|
||||
|
||||
/// Progressive merge state
|
||||
#[derive(Debug)]
|
||||
pub struct ProgressiveMergeState {
|
||||
/// Current step in the merge process
|
||||
pub current_step: usize,
|
||||
|
||||
/// Current merged parameters
|
||||
pub current_parameters: IndexMap<String, ParameterTensor>,
|
||||
|
||||
/// History of checkpoints
|
||||
pub checkpoints: Vec<MergeCheckpoint>,
|
||||
|
||||
/// Validation metrics
|
||||
pub validation_metrics: Vec<f32>,
|
||||
}
|
||||
|
||||
/// Checkpoint information
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MergeCheckpoint {
|
||||
/// Step number
|
||||
pub step: usize,
|
||||
|
||||
/// Parameters at this checkpoint
|
||||
pub parameters: Arc<IndexMap<String, ParameterTensor>>,
|
||||
|
||||
/// Validation score if available
|
||||
pub validation_score: Option<f32>,
|
||||
}
|
||||
|
||||
/// Progressive merge implementation
|
||||
pub struct ProgressiveMerge {
|
||||
config: ProgressiveConfig,
|
||||
state: Option<ProgressiveMergeState>,
|
||||
}
|
||||
|
||||
impl ProgressiveMerge {
|
||||
/// Create a new progressive merge instance
|
||||
pub fn new(config: ProgressiveConfig) -> Self {
|
||||
Self {
|
||||
config,
|
||||
state: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Initialize the merge process
|
||||
pub fn initialize(&mut self, base_model: Model, target_models: Vec<Model>) -> Result<()> {
|
||||
if target_models.is_empty() {
|
||||
return Err(MergeError::config(
|
||||
"target_models",
|
||||
"Progressive merge requires at least one target model",
|
||||
));
|
||||
}
|
||||
|
||||
self.state = Some(ProgressiveMergeState {
|
||||
current_step: 0,
|
||||
current_parameters: base_model.parameters.clone(),
|
||||
checkpoints: Vec::new(),
|
||||
validation_metrics: Vec::new(),
|
||||
});
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Perform one step of progressive merging
|
||||
pub fn merge_step(&mut self, target_models: &[Model]) -> Result<bool> {
|
||||
let state = self
|
||||
.state
|
||||
.as_mut()
|
||||
.ok_or_else(|| MergeError::internal("Progressive merge not initialized"))?;
|
||||
|
||||
if state.current_step >= self.config.num_steps {
|
||||
return Ok(true); // Merge complete
|
||||
}
|
||||
|
||||
// Calculate blend weight for this step
|
||||
let blend_weight =
|
||||
Self::calculate_blend_weight_static(state.current_step, self.config.num_steps);
|
||||
|
||||
// Merge parameters progressively
|
||||
let mut updated_parameters = IndexMap::new();
|
||||
for (param_name, current_param) in &state.current_parameters {
|
||||
let merged_param = Self::merge_parameter_progressive_static(
|
||||
current_param,
|
||||
target_models,
|
||||
param_name,
|
||||
blend_weight,
|
||||
)?;
|
||||
updated_parameters.insert(param_name.clone(), merged_param);
|
||||
}
|
||||
|
||||
state.current_parameters = updated_parameters;
|
||||
state.current_step += 1;
|
||||
|
||||
// Create checkpoint if needed
|
||||
if state.current_step % self.config.validation_frequency.max(1) == 0 {
|
||||
let checkpoint = MergeCheckpoint {
|
||||
step: state.current_step,
|
||||
parameters: Arc::new(state.current_parameters.clone()),
|
||||
validation_score: None,
|
||||
};
|
||||
state.checkpoints.push(checkpoint);
|
||||
}
|
||||
|
||||
// Check early stopping
|
||||
if self.config.early_stopping_patience.is_some() {
|
||||
let threshold = self.config.metric_threshold;
|
||||
if let Some(last_metric) = state.validation_metrics.last()
|
||||
&& *last_metric >= threshold
|
||||
{
|
||||
return Ok(true); // Early stopping
|
||||
}
|
||||
}
|
||||
|
||||
Ok(state.current_step >= self.config.num_steps)
|
||||
}
|
||||
|
||||
/// Calculate blend weight for a given step
|
||||
fn calculate_blend_weight(&self, step: usize) -> f32 {
|
||||
Self::calculate_blend_weight_static(step, self.config.num_steps)
|
||||
}
|
||||
|
||||
/// Calculate blend weight for a given step (static version)
|
||||
fn calculate_blend_weight_static(step: usize, num_steps: usize) -> f32 {
|
||||
// Simple linear interpolation
|
||||
step as f32 / num_steps as f32
|
||||
}
|
||||
|
||||
/// Merge a single parameter progressively (static version)
|
||||
fn merge_parameter_progressive_static(
|
||||
current_param: &ParameterTensor,
|
||||
target_models: &[Model],
|
||||
param_name: &str,
|
||||
blend_weight: f32,
|
||||
) -> Result<ParameterTensor> {
|
||||
let mut merged = current_param.clone();
|
||||
|
||||
// Collect target parameters
|
||||
let mut target_params = Vec::new();
|
||||
for model in target_models {
|
||||
if let Some(param) = model.parameters.get(param_name)
|
||||
&& param.shape == current_param.shape
|
||||
{
|
||||
target_params.push(param);
|
||||
}
|
||||
}
|
||||
|
||||
if target_params.is_empty() {
|
||||
return Ok(merged); // No matching parameters to merge
|
||||
}
|
||||
|
||||
// Progressive blending
|
||||
let keep_weight = 1.0 - blend_weight;
|
||||
let target_weight = blend_weight / target_params.len() as f32;
|
||||
|
||||
for i in 0..merged.data.len() {
|
||||
// Scale current value
|
||||
merged.data[i] *= keep_weight;
|
||||
|
||||
// Add target contributions
|
||||
for target_param in &target_params {
|
||||
if i < target_param.data.len() {
|
||||
merged.data[i] += target_param.data[i] * target_weight;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(merged)
|
||||
}
|
||||
|
||||
/// Create a checkpoint of the current state
|
||||
fn create_checkpoint(&mut self) -> Result<()> {
|
||||
let state = self
|
||||
.state
|
||||
.as_mut()
|
||||
.ok_or_else(|| MergeError::internal("Progressive merge not initialized"))?;
|
||||
|
||||
let checkpoint = MergeCheckpoint {
|
||||
step: state.current_step,
|
||||
parameters: Arc::new(state.current_parameters.clone()),
|
||||
validation_score: None, // Would be set by external validation
|
||||
};
|
||||
|
||||
state.checkpoints.push(checkpoint);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Complete the merge and return the final model
|
||||
pub fn finalize(self) -> Result<MergedModel> {
|
||||
let state = self
|
||||
.state
|
||||
.ok_or_else(|| MergeError::internal("Progressive merge not initialized"))?;
|
||||
|
||||
let merged_model = Model {
|
||||
id: Uuid::new_v4(),
|
||||
name: format!("progressive_merged_{}", chrono::Utc::now().timestamp()),
|
||||
architecture: crate::types::ModelArchitecture {
|
||||
arch_type: "progressive_merge".to_string(),
|
||||
num_layers: 0,
|
||||
hidden_dim: 0,
|
||||
params: std::collections::HashMap::new(),
|
||||
},
|
||||
parameters: state.current_parameters,
|
||||
metadata: crate::types::ModelMetadata::default(),
|
||||
config: crate::types::ModelConfig::default(),
|
||||
};
|
||||
|
||||
Ok(MergedModel {
|
||||
model: merged_model,
|
||||
merge_info: MergeInfo {
|
||||
strategy: "progressive".to_string(),
|
||||
source_models: Vec::new(), // Would be tracked during initialization
|
||||
merged_at: chrono::Utc::now(),
|
||||
config: serde_json::to_value(&self.config).unwrap_or(serde_json::Value::Null),
|
||||
statistics: MergeStatistics {
|
||||
parameters_merged: 0,
|
||||
parameters_conflicted: 0,
|
||||
parameters_dropped: 0,
|
||||
duration_ms: 0,
|
||||
memory_usage_mb: 0,
|
||||
gpu_memory_mb: Some(0),
|
||||
},
|
||||
},
|
||||
quality_metrics: QualityMetrics {
|
||||
consistency_score: 1.0,
|
||||
complexity_score: 1.0,
|
||||
quality_indicators: std::collections::HashMap::new(),
|
||||
validation_results: std::collections::HashMap::new(),
|
||||
performance_predictions: std::collections::HashMap::new(),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// Update validation metric for the current step
|
||||
pub fn update_validation_metric(&mut self, metric: f32) -> Result<()> {
|
||||
let state = self
|
||||
.state
|
||||
.as_mut()
|
||||
.ok_or_else(|| MergeError::internal("Progressive merge not initialized"))?;
|
||||
|
||||
state.validation_metrics.push(metric);
|
||||
|
||||
// Update checkpoint if exists
|
||||
if let Some(checkpoint) = state.checkpoints.last_mut()
|
||||
&& checkpoint.step == state.current_step
|
||||
{
|
||||
checkpoint.validation_score = Some(metric);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Rollback to a previous checkpoint
|
||||
pub fn rollback_to_checkpoint(&mut self, checkpoint_index: usize) -> Result<()> {
|
||||
let state = self
|
||||
.state
|
||||
.as_mut()
|
||||
.ok_or_else(|| MergeError::internal("Progressive merge not initialized"))?;
|
||||
|
||||
if checkpoint_index >= state.checkpoints.len() {
|
||||
return Err(MergeError::config(
|
||||
"checkpoint_index",
|
||||
"Invalid checkpoint index",
|
||||
));
|
||||
}
|
||||
|
||||
let checkpoint = &state.checkpoints[checkpoint_index];
|
||||
state.current_parameters = (*checkpoint.parameters).clone();
|
||||
state.current_step = checkpoint.step;
|
||||
|
||||
// Remove later checkpoints
|
||||
state.checkpoints.truncate(checkpoint_index + 1);
|
||||
state.validation_metrics.truncate(checkpoint_index + 1);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Convenience function for Progressive merging
|
||||
pub async fn merge_models(models: &[Model], config: &ProgressiveConfig) -> Result<MergedModel> {
|
||||
let mut merger = ProgressiveMerge::new(config.clone());
|
||||
|
||||
// Initialize with first model as base, rest as targets
|
||||
if models.is_empty() {
|
||||
return Err(MergeError::config(
|
||||
"models",
|
||||
"Progressive merge requires at least one model",
|
||||
));
|
||||
}
|
||||
|
||||
let base_model = models[0].clone();
|
||||
let target_models = models[1..].to_vec();
|
||||
|
||||
merger.initialize(base_model, target_models.clone())?;
|
||||
|
||||
// Perform all merge steps
|
||||
loop {
|
||||
let done = merger.merge_step(&target_models)?;
|
||||
if done {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
merger.finalize()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_progressive_merge_creation() {
|
||||
let config = ProgressiveConfig::default();
|
||||
let merge = ProgressiveMerge::new(config);
|
||||
assert!(merge.state.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_blend_weight_calculation() {
|
||||
let config = ProgressiveConfig {
|
||||
num_steps: 10,
|
||||
validation_frequency: 1,
|
||||
early_stopping_patience: None,
|
||||
metric_threshold: 0.01,
|
||||
enable_rollback: true,
|
||||
};
|
||||
let merge = ProgressiveMerge::new(config);
|
||||
|
||||
assert_eq!(merge.calculate_blend_weight(0), 0.0);
|
||||
assert_eq!(merge.calculate_blend_weight(5), 0.5);
|
||||
assert_eq!(merge.calculate_blend_weight(10), 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exponential_blend_weight() {
|
||||
let config = ProgressiveConfig {
|
||||
num_steps: 10,
|
||||
validation_frequency: 1,
|
||||
early_stopping_patience: Some(3),
|
||||
metric_threshold: 0.01,
|
||||
enable_rollback: true,
|
||||
};
|
||||
let merge = ProgressiveMerge::new(config);
|
||||
|
||||
let weight = merge.calculate_blend_weight(5);
|
||||
assert!(weight > 0.0 && weight < 1.0);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user