Files
rustytorch/crates/training/rtx-model-merging/src/lib.rs
T
2026-03-04 00:08:42 +00:00

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());
}
}