234 lines
7.0 KiB
Rust
234 lines
7.0 KiB
Rust
//! # 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<T> = Result<T, OptimizationError>;
|
|
|
|
/// 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<f64>,
|
|
/// Quality degradation (0.0 = none, 1.0 = complete)
|
|
pub quality_loss: Option<f64>,
|
|
/// Strategies applied
|
|
pub strategies: Vec<OptimizationStrategy>,
|
|
}
|
|
|
|
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<OptimizationStrategy> {
|
|
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);
|
|
}
|
|
}
|