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

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