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,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(&param2.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(&param1, &param2)
.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);
}
}