pub mod distillation;
pub mod graph_opt;
pub mod pruning;
pub mod quantization;
pub use distillation::{
DenseLayer,
DistillationConfig,
DistillationConfigBuilder,
DistillationLoss,
DistillationStats,
DistillationTrainer,
EarlyStopping,
ForwardCache,
LearningRateSchedule,
MLPGradients,
OptimizerType,
SimpleMLP,
SimpleRng,
Temperature,
TrainingState,
cross_entropy_loss,
cross_entropy_with_label,
kl_divergence,
kl_divergence_from_logits,
log_softmax,
mse_loss,
soft_targets,
softmax,
train_student_model,
};
pub use graph_opt::{
GraphOptConfig, OptimizationBenchmark, apply_graph_optimization,
apply_graph_optimization_from_bytes, benchmark_optimization,
};
pub use pruning::{
FineTuneCallback,
GradientInfo,
ImportanceMethod,
LotteryTicketState,
MaskCreationMode,
NoOpFineTune,
PruningConfig,
PruningConfigBuilder,
PruningGranularity,
PruningMask,
PruningSchedule,
PruningStats,
PruningStrategy,
UnstructuredPruner,
WeightStatistics,
WeightTensor,
compute_channel_importance,
compute_gradient_importance,
compute_magnitude_importance,
compute_taylor_importance,
iterative_pruning,
prune_model,
prune_weights_direct,
prune_weights_with_gradients,
select_weights_to_prune,
structured_pruning,
unstructured_pruning,
};
pub use quantization::{
QuantizationConfig, QuantizationMode, QuantizationParams, QuantizationResult, QuantizationType,
calibrate_quantization, dequantize_tensor, quantize_model, quantize_tensor,
};
use crate::error::Result;
use std::path::Path;
use tracing::info;
#[derive(Debug, Clone)]
pub struct OptimizationStats {
pub original_size: usize,
pub optimized_size: usize,
pub compression_ratio: f32,
pub speedup: f32,
pub accuracy_delta: f32,
}
impl OptimizationStats {
#[must_use]
pub fn new(
original_size: usize,
optimized_size: usize,
speedup: f32,
accuracy_delta: f32,
) -> Self {
let compression_ratio = if optimized_size > 0 {
original_size as f32 / optimized_size as f32
} else {
0.0
};
Self {
original_size,
optimized_size,
compression_ratio,
speedup,
accuracy_delta,
}
}
#[must_use]
pub fn size_reduction(&self) -> usize {
self.original_size.saturating_sub(self.optimized_size)
}
#[must_use]
pub fn size_reduction_percent(&self) -> f32 {
if self.original_size > 0 {
(self.size_reduction() as f32 / self.original_size as f32) * 100.0
} else {
0.0
}
}
#[must_use]
pub fn is_worthwhile(&self) -> bool {
self.size_reduction_percent() > 20.0 && self.accuracy_delta.abs() < 2.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OptimizationProfile {
Accuracy,
Balanced,
Speed,
Size,
}
pub struct OptimizationPipeline {
pub quantization: Option<QuantizationConfig>,
pub pruning: Option<PruningConfig>,
pub weight_sharing: bool,
pub operator_fusion: bool,
pub graph_opt_config: Option<GraphOptConfig>,
}
impl OptimizationPipeline {
#[must_use]
pub fn effective_graph_opt_config(&self) -> GraphOptConfig {
if let Some(ref config) = self.graph_opt_config {
config.clone()
} else if self.operator_fusion {
GraphOptConfig::default()
} else {
GraphOptConfig::none()
}
}
#[must_use]
pub fn opt_level(&self) -> oxionnx::OptLevel {
self.effective_graph_opt_config().to_opt_level()
}
}
impl OptimizationPipeline {
#[must_use]
pub fn from_profile(profile: OptimizationProfile) -> Self {
match profile {
OptimizationProfile::Accuracy => Self {
quantization: Some(
QuantizationConfig::builder()
.quantization_type(QuantizationType::Float16)
.build(),
),
pruning: None,
weight_sharing: false,
operator_fusion: true,
graph_opt_config: Some(GraphOptConfig::default()),
},
OptimizationProfile::Balanced => Self {
quantization: Some(
QuantizationConfig::builder()
.quantization_type(QuantizationType::Int8)
.per_channel(true)
.build(),
),
pruning: Some(
PruningConfig::builder()
.sparsity_target(0.3)
.strategy(PruningStrategy::Magnitude)
.build(),
),
weight_sharing: true,
operator_fusion: true,
graph_opt_config: Some(GraphOptConfig::default()),
},
OptimizationProfile::Speed => Self {
quantization: Some(
QuantizationConfig::builder()
.quantization_type(QuantizationType::Int8)
.per_channel(true)
.build(),
),
pruning: Some(
PruningConfig::builder()
.sparsity_target(0.5)
.strategy(PruningStrategy::Structured)
.build(),
),
weight_sharing: true,
operator_fusion: true,
graph_opt_config: Some(GraphOptConfig {
constant_folding: true,
dead_node_elimination: true,
common_subexpression_elimination: true,
operator_fusion: true,
}),
},
OptimizationProfile::Size => Self {
quantization: Some(
QuantizationConfig::builder()
.quantization_type(QuantizationType::Int8)
.per_channel(true)
.build(),
),
pruning: Some(
PruningConfig::builder()
.sparsity_target(0.7)
.strategy(PruningStrategy::Structured)
.build(),
),
weight_sharing: true,
operator_fusion: true,
graph_opt_config: Some(GraphOptConfig::default()),
},
}
}
pub fn optimize<P: AsRef<std::path::Path>>(
&self,
input_path: P,
output_path: P,
) -> Result<OptimizationStats> {
use tracing::info;
info!("Running optimization pipeline");
let input = input_path.as_ref();
let output = output_path.as_ref();
let original_size = std::fs::metadata(input)
.map(|m| m.len() as usize)
.unwrap_or(0);
let mut current_path = input.to_path_buf();
if let Some(ref config) = self.pruning {
let pruned_path = output.with_extension("pruned.onnx");
prune_model(¤t_path, &pruned_path, config)?;
current_path = pruned_path;
}
if let Some(ref config) = self.quantization {
let quantized_path = output.with_extension("quantized.onnx");
quantize_model(¤t_path, &quantized_path, config)?;
current_path = quantized_path;
}
std::fs::rename(¤t_path, output)?;
let optimized_size = std::fs::metadata(output)
.map(|m| m.len() as usize)
.unwrap_or(0);
let opt_level = self.opt_level();
let speedup = Self::measure_speedup(input, output, opt_level)?;
let accuracy_delta = Self::estimate_accuracy_delta(self);
Ok(OptimizationStats::new(
original_size,
optimized_size,
speedup,
accuracy_delta,
))
}
fn measure_speedup(
original_path: &Path,
optimized_path: &Path,
opt_level: oxionnx::OptLevel,
) -> Result<f32> {
const WARMUP_ITERS: usize = 5;
const BENCH_ITERS: usize = 20;
if !original_path.exists() || !optimized_path.exists() {
info!("Skipping speedup measurement: model files not accessible");
return Ok(1.5); }
let dummy_input = vec![0.0f32; 224 * 224 * 3]; let input_shape = vec![1, 3, 224, 224];
let original_time = match Self::benchmark_model(
original_path,
&dummy_input,
&input_shape,
WARMUP_ITERS,
BENCH_ITERS,
oxionnx::OptLevel::None,
) {
Ok(t) => t,
Err(e) => {
info!("Could not benchmark original model: {}, using estimate", e);
return Ok(1.5);
}
};
let optimized_time = match Self::benchmark_model(
optimized_path,
&dummy_input,
&input_shape,
WARMUP_ITERS,
BENCH_ITERS,
opt_level,
) {
Ok(t) => t,
Err(e) => {
info!("Could not benchmark optimized model: {}, using estimate", e);
return Ok(1.5);
}
};
if optimized_time > 0.0 {
let speedup = (original_time / optimized_time) as f32;
info!(
"Measured speedup: {:.2}x (original: {:.2}ms, optimized: {:.2}ms)",
speedup,
original_time * 1000.0,
optimized_time * 1000.0
);
Ok(speedup)
} else {
Ok(1.5) }
}
fn benchmark_model(
model_path: &Path,
input: &[f32],
input_shape: &[usize],
warmup_iters: usize,
bench_iters: usize,
opt_level: oxionnx::OptLevel,
) -> Result<f64> {
use ndarray::{Array, IxDyn};
use oxionnx::{SessionBuilder, Tensor};
use std::time::Instant;
let session = SessionBuilder::new()
.with_optimization_level(opt_level)
.commit_from_file(model_path)
.map_err(|e| crate::error::ModelError::LoadFailed {
reason: format!("Failed to load model for benchmarking: {}", e),
})?;
let input_name = session
.input_info()
.first()
.ok_or_else(|| crate::error::ModelError::LoadFailed {
reason: "No input tensors found in model".to_string(),
})?
.name
.clone();
let array_shape: Vec<usize> = input_shape.to_vec();
let total_elements: usize = array_shape.iter().product();
if input.len() != total_elements {
return Err(crate::error::InferenceError::InvalidInputShape {
expected: array_shape.clone(),
actual: vec![input.len()],
}
.into());
}
let input_array =
Array::from_shape_vec(IxDyn(&array_shape), input.to_vec()).map_err(|e| {
crate::error::InferenceError::Failed {
reason: format!("Failed to create input array: {}", e),
}
})?;
for _ in 0..warmup_iters {
let input_tensor = Tensor::from_ndarray_view(input_array.view());
let inputs_map =
oxionnx::inputs![input_name.as_str() => input_tensor].map_err(|e| {
crate::error::InferenceError::Failed {
reason: format!("Failed to build inputs map: {}", e),
}
})?;
let _ = session
.run(&inputs_map)
.map_err(|e| crate::error::InferenceError::Failed {
reason: format!("Warmup inference failed: {}", e),
})?;
}
let start = Instant::now();
for _ in 0..bench_iters {
let input_tensor = Tensor::from_ndarray_view(input_array.view());
let inputs_map =
oxionnx::inputs![input_name.as_str() => input_tensor].map_err(|e| {
crate::error::InferenceError::Failed {
reason: format!("Failed to build inputs map: {}", e),
}
})?;
let _ = session
.run(&inputs_map)
.map_err(|e| crate::error::InferenceError::Failed {
reason: format!("Benchmark inference failed: {}", e),
})?;
}
let elapsed = start.elapsed();
let avg_time = elapsed.as_secs_f64() / bench_iters as f64;
Ok(avg_time)
}
fn estimate_accuracy_delta(&self) -> f32 {
let mut delta = 0.0f32;
if let Some(ref quant) = self.quantization {
delta += match quant.quantization_type {
QuantizationType::Float16 => -0.1, QuantizationType::Int8 => -0.5, QuantizationType::UInt8 => -0.5, QuantizationType::Int4 => -2.0, };
}
if let Some(ref prune) = self.pruning {
delta += -prune.sparsity_target * 2.0; }
delta
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_optimization_stats() {
let stats = OptimizationStats::new(
1000000, 250000, 2.0, -0.5, );
assert_eq!(stats.size_reduction(), 750000);
assert!((stats.size_reduction_percent() - 75.0).abs() < 0.1);
assert!((stats.compression_ratio - 4.0).abs() < 0.1);
assert!(stats.is_worthwhile());
}
#[test]
fn test_optimization_profile_accuracy() {
let pipeline = OptimizationPipeline::from_profile(OptimizationProfile::Accuracy);
assert!(pipeline.quantization.is_some());
assert!(pipeline.pruning.is_none());
assert!(pipeline.operator_fusion);
assert!(pipeline.graph_opt_config.is_some());
assert_eq!(pipeline.opt_level(), oxionnx::OptLevel::All);
}
#[test]
fn test_optimization_profile_speed() {
let pipeline = OptimizationPipeline::from_profile(OptimizationProfile::Speed);
assert!(pipeline.quantization.is_some());
assert!(pipeline.pruning.is_some());
assert!(pipeline.weight_sharing);
assert!(pipeline.operator_fusion);
assert_eq!(pipeline.opt_level(), oxionnx::OptLevel::All);
}
#[test]
fn test_optimization_profile_size() {
let pipeline = OptimizationPipeline::from_profile(OptimizationProfile::Size);
assert!(pipeline.quantization.is_some());
assert!(pipeline.pruning.is_some());
if let Some(pruning) = &pipeline.pruning {
assert!(pruning.sparsity_target >= 0.6);
}
assert_eq!(pipeline.opt_level(), oxionnx::OptLevel::All);
}
#[test]
fn test_effective_graph_opt_config_from_fusion_flag() {
let pipeline = OptimizationPipeline {
quantization: None,
pruning: None,
weight_sharing: false,
operator_fusion: true,
graph_opt_config: None,
};
let config = pipeline.effective_graph_opt_config();
assert!(config.any_enabled());
assert_eq!(pipeline.opt_level(), oxionnx::OptLevel::All);
let pipeline_no_fusion = OptimizationPipeline {
quantization: None,
pruning: None,
weight_sharing: false,
operator_fusion: false,
graph_opt_config: None,
};
let config_none = pipeline_no_fusion.effective_graph_opt_config();
assert!(!config_none.any_enabled());
assert_eq!(pipeline_no_fusion.opt_level(), oxionnx::OptLevel::None);
}
#[test]
fn test_effective_graph_opt_config_explicit_override() {
let custom_config = GraphOptConfig {
constant_folding: true,
dead_node_elimination: false,
common_subexpression_elimination: false,
operator_fusion: false,
};
let pipeline = OptimizationPipeline {
quantization: None,
pruning: None,
weight_sharing: false,
operator_fusion: false, graph_opt_config: Some(custom_config),
};
let config = pipeline.effective_graph_opt_config();
assert!(config.constant_folding);
assert!(!config.operator_fusion);
assert_eq!(pipeline.opt_level(), oxionnx::OptLevel::All);
}
}