184 lines
5.6 KiB
Rust
184 lines
5.6 KiB
Rust
//! # RTX Model Merging
|
|
//!
|
|
//! Advanced model merging algorithms and strategies for RustyTorch.
|
|
//!
|
|
//! This crate provides comprehensive model merging capabilities including:
|
|
//! - TIES (Task-specific Interference Elimination)
|
|
//! - DARE (Drop And REscale) merging
|
|
//! - SLERP (Spherical Linear Interpolation)
|
|
//! - Model soups and ensemble methods
|
|
//! - Frankenmerging techniques
|
|
//! - Dynamic expert selection and routing
|
|
//! - Merge conflict resolution
|
|
//! - Performance evaluation and benchmarking
|
|
//!
|
|
//! # Features
|
|
//!
|
|
//! - **Mathematical Precision**: Exact implementations from research papers
|
|
//! - **Memory Efficiency**: Streaming parameter processing for large models
|
|
//! - **Performance Optimization**: Parallel processing and GPU acceleration
|
|
//! - **Robustness**: Comprehensive error handling and validation
|
|
//! - **Configuration**: Flexible YAML/TOML configuration system
|
|
//!
|
|
//! # Example
|
|
//!
|
|
//! ```rust
|
|
//! use rtx_model_merging::{MergeStrategy, ModelMerger, TiesConfig};
|
|
//!
|
|
//! # async fn example() -> anyhow::Result<()> {
|
|
//! let merger = ModelMerger::new()?;
|
|
//! let strategy = MergeStrategy::Ties(TiesConfig::default());
|
|
//!
|
|
//! let merged_model = merger
|
|
//! .merge_models(&["model1.safetensors", "model2.safetensors"], &strategy)
|
|
//! .await?;
|
|
//! # Ok(())
|
|
//! # }
|
|
//! ```
|
|
|
|
pub mod algorithms;
|
|
pub mod config;
|
|
pub mod conflict_resolution;
|
|
pub mod error;
|
|
pub mod evaluation;
|
|
pub mod expert_selection;
|
|
pub mod model_loader;
|
|
pub mod planning;
|
|
pub mod strategies;
|
|
pub mod types;
|
|
|
|
// Re-export core types and functions for convenience
|
|
pub use algorithms::*;
|
|
pub use config::*;
|
|
pub use conflict_resolution::*;
|
|
pub use error::{MergeError, Result};
|
|
pub use evaluation::*;
|
|
pub use expert_selection::*;
|
|
pub use model_loader::*;
|
|
pub use planning::*;
|
|
pub use strategies::*;
|
|
pub use types::*;
|
|
|
|
use std::path::Path;
|
|
use std::sync::Arc;
|
|
|
|
/// Main model merging orchestrator
|
|
pub struct ModelMerger {
|
|
loader: Arc<ModelLoader>,
|
|
validator: Arc<MergeValidator>,
|
|
evaluator: Arc<PerformanceEvaluator>,
|
|
}
|
|
|
|
impl ModelMerger {
|
|
/// Create a new model merger with default configuration
|
|
pub fn new() -> Result<Self> {
|
|
let loader = Arc::new(ModelLoader::new()?);
|
|
let validator = Arc::new(MergeValidator::new());
|
|
let evaluator = Arc::new(PerformanceEvaluator::new()?);
|
|
|
|
Ok(Self {
|
|
loader,
|
|
validator,
|
|
evaluator,
|
|
})
|
|
}
|
|
|
|
/// Create a new model merger with custom configuration
|
|
pub fn with_config(config: MergeConfig) -> Result<Self> {
|
|
let loader = Arc::new(ModelLoader::with_config(config.loader)?);
|
|
let validator = Arc::new(MergeValidator::with_config(config.validation));
|
|
let evaluator = Arc::new(PerformanceEvaluator::with_config(config.evaluation)?);
|
|
|
|
Ok(Self {
|
|
loader,
|
|
validator,
|
|
evaluator,
|
|
})
|
|
}
|
|
|
|
/// Merge multiple models using the specified strategy
|
|
pub async fn merge_models<P: AsRef<Path>>(
|
|
&self,
|
|
model_paths: &[P],
|
|
strategy: &MergeStrategy,
|
|
) -> Result<MergedModel> {
|
|
tracing::info!("Starting model merge with strategy: {:?}", strategy);
|
|
|
|
// Load all models
|
|
let models = self.loader.load_models(model_paths).await?;
|
|
|
|
// Validate compatibility
|
|
self.validator.validate_compatibility(&models)?;
|
|
|
|
// Execute merge strategy
|
|
let merged = match strategy {
|
|
MergeStrategy::Ties(config) => algorithms::ties::merge_models(&models, config).await?,
|
|
MergeStrategy::Dare(config) => algorithms::dare::merge_models(&models, config).await?,
|
|
MergeStrategy::Slerp(config) => {
|
|
algorithms::slerp::merge_models(&models, config).await?
|
|
}
|
|
MergeStrategy::TaskArithmetic(config) => {
|
|
algorithms::task_arithmetic::merge_models(&models, config).await?
|
|
}
|
|
MergeStrategy::Fisher(config) => {
|
|
algorithms::fisher::merge_models(&models, config).await?
|
|
}
|
|
MergeStrategy::ModelSoup(config) => {
|
|
strategies::model_soup::merge_models(&models, config).await?
|
|
}
|
|
MergeStrategy::Frankenmerge(config) => {
|
|
strategies::frankenmerge::merge_models(&models, config).await?
|
|
}
|
|
MergeStrategy::Progressive(config) => {
|
|
strategies::progressive::merge_models(&models, config).await?
|
|
}
|
|
};
|
|
|
|
// Validate the merged model
|
|
self.validator.validate_merged_model(&merged)?;
|
|
|
|
tracing::info!("Model merge completed successfully");
|
|
Ok(merged)
|
|
}
|
|
|
|
/// Evaluate the performance of a merged model
|
|
pub async fn evaluate_model(&self, model: &MergedModel) -> Result<EvaluationResults> {
|
|
self.evaluator.evaluate(model).await
|
|
}
|
|
|
|
/// Recommend the best merge strategy for given models
|
|
pub async fn recommend_strategy<P: AsRef<Path>>(
|
|
&self,
|
|
model_paths: &[P],
|
|
) -> Result<MergeStrategy> {
|
|
let models = self.loader.load_model_metadata(model_paths).await?;
|
|
let planner = MergePlanner::new();
|
|
planner.recommend_strategy(&models)
|
|
}
|
|
}
|
|
|
|
impl Default for ModelMerger {
|
|
fn default() -> Self {
|
|
Self::new().expect("Failed to create default ModelMerger")
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use tempfile::TempDir;
|
|
|
|
#[tokio::test]
|
|
async fn test_model_merger_creation() {
|
|
let merger = ModelMerger::new();
|
|
assert!(merger.is_ok());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_model_merger_with_config() {
|
|
let config = MergeConfig::default();
|
|
let merger = ModelMerger::with_config(config);
|
|
assert!(merger.is_ok());
|
|
}
|
|
}
|