//! # RTX Model Optimization Core //! //! Unified framework for model optimization in RustyTorch, consolidating: //! - **Compression**: Quantization, pruning, distillation, LoRA //! - **Merging**: Model merging algorithms (TIES, DARE, SLERP) //! //! ## Features //! //! Enable optimization techniques via Cargo features: //! - `compress` - Model compression (quantization, pruning, distillation) (default) //! - `merging` - Model merging algorithms (default) //! - `cuda` - CUDA GPU acceleration for compression //! - `metal` - Metal GPU acceleration for merging //! - `all` - All features //! //! ## Example //! //! ```rust,ignore //! use rtx_optim_core::{OptimizationPipeline, OptimizationStrategy}; //! //! // Create an optimization pipeline //! let pipeline = OptimizationPipeline::new() //! .with_quantization(8) // INT8 quantization //! .with_pruning(0.3) // 30% sparsity //! .build(); //! //! // Optimize a model //! let optimized = pipeline.optimize(&model)?; //! ``` #![forbid(unsafe_code)] use serde::{Deserialize, Serialize}; use thiserror::Error; // Re-export compression module #[cfg(feature = "compress")] pub use rtx_compress as compress; // Re-export merging module #[cfg(feature = "merging")] pub use rtx_model_merging as merging; /// Model optimization strategies #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum OptimizationStrategy { /// Quantization - reduce precision (INT8, INT4, etc.) Quantization, /// Pruning - remove unimportant weights Pruning, /// Distillation - train smaller model from larger Distillation, /// LoRA - low-rank adaptation LoRA, /// Model merging - combine multiple models Merging, /// KV Cache compression KVCacheCompression, } impl std::fmt::Display for OptimizationStrategy { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { OptimizationStrategy::Quantization => write!(f, "Quantization"), OptimizationStrategy::Pruning => write!(f, "Pruning"), OptimizationStrategy::Distillation => write!(f, "Distillation"), OptimizationStrategy::LoRA => write!(f, "LoRA"), OptimizationStrategy::Merging => write!(f, "Model Merging"), OptimizationStrategy::KVCacheCompression => write!(f, "KV Cache Compression"), } } } /// Optimization error types #[derive(Debug, Error)] pub enum OptimizationError { #[error("Compression error: {message}")] CompressionError { message: String }, #[error("Merging error: {message}")] MergingError { message: String }, #[error("Configuration error: {message}")] ConfigError { message: String }, #[error("Unsupported strategy: {strategy}")] UnsupportedStrategy { strategy: String }, #[error("Model incompatible with optimization: {reason}")] IncompatibleModel { reason: String }, #[error("IO error: {0}")] Io(#[from] std::io::Error), } /// Result type for optimization operations pub type OptimResult = Result; /// Optimization metrics #[derive(Debug, Clone, Serialize, Deserialize)] pub struct OptimizationMetrics { /// Original model size in bytes pub original_size: u64, /// Optimized model size in bytes pub optimized_size: u64, /// Compression ratio (original / optimized) pub compression_ratio: f64, /// Speedup factor (if measured) pub speedup: Option, /// Quality degradation (0.0 = none, 1.0 = complete) pub quality_loss: Option, /// Strategies applied pub strategies: Vec, } impl OptimizationMetrics { /// Create new metrics pub fn new(original_size: u64, optimized_size: u64) -> Self { let compression_ratio = if optimized_size > 0 { original_size as f64 / optimized_size as f64 } else { 0.0 }; Self { original_size, optimized_size, compression_ratio, speedup: None, quality_loss: None, strategies: Vec::new(), } } /// Add a strategy to the metrics pub fn with_strategy(mut self, strategy: OptimizationStrategy) -> Self { self.strategies.push(strategy); self } /// Set speedup factor pub fn with_speedup(mut self, speedup: f64) -> Self { self.speedup = Some(speedup); self } /// Set quality loss pub fn with_quality_loss(mut self, loss: f64) -> Self { self.quality_loss = Some(loss); self } } /// Check if a feature is enabled pub fn is_feature_enabled(strategy: OptimizationStrategy) -> bool { match strategy { OptimizationStrategy::Quantization | OptimizationStrategy::Pruning | OptimizationStrategy::Distillation | OptimizationStrategy::LoRA | OptimizationStrategy::KVCacheCompression => cfg!(feature = "compress"), OptimizationStrategy::Merging => cfg!(feature = "merging"), } } /// Get list of enabled optimization strategies pub fn enabled_strategies() -> Vec { let mut strategies = Vec::new(); #[cfg(feature = "compress")] { strategies.push(OptimizationStrategy::Quantization); strategies.push(OptimizationStrategy::Pruning); strategies.push(OptimizationStrategy::Distillation); strategies.push(OptimizationStrategy::LoRA); strategies.push(OptimizationStrategy::KVCacheCompression); } #[cfg(feature = "merging")] { strategies.push(OptimizationStrategy::Merging); } strategies } /// Prelude module for common imports pub mod prelude { pub use super::{OptimResult, OptimizationError, OptimizationMetrics, OptimizationStrategy}; #[cfg(feature = "compress")] pub use rtx_compress::{CompressionConfig, CompressionPipeline}; #[cfg(feature = "merging")] pub use rtx_model_merging::{MergeConfig, MergeStrategy}; } #[cfg(test)] mod tests { use super::*; #[test] fn test_strategy_display() { assert_eq!( OptimizationStrategy::Quantization.to_string(), "Quantization" ); assert_eq!(OptimizationStrategy::Pruning.to_string(), "Pruning"); assert_eq!(OptimizationStrategy::Merging.to_string(), "Model Merging"); } #[test] fn test_optimization_metrics() { let metrics = OptimizationMetrics::new(1000, 250) .with_strategy(OptimizationStrategy::Quantization) .with_speedup(2.5) .with_quality_loss(0.01); assert_eq!(metrics.original_size, 1000); assert_eq!(metrics.optimized_size, 250); assert!((metrics.compression_ratio - 4.0).abs() < 0.001); assert_eq!(metrics.speedup, Some(2.5)); assert_eq!(metrics.quality_loss, Some(0.01)); } #[test] fn test_enabled_strategies() { let strategies = enabled_strategies(); // With default features, both compress and merging should be enabled #[cfg(all(feature = "compress", feature = "merging"))] assert!(strategies.len() >= 6); } }