use candle_core::{DType, Device};
use mlmf::{
load_model, ActivationStats, LoadOptions, QuantizationConfig, QuantizationContext,
QuantizationEngine, QuantizationScheme, QuantizationType,
};
use std::collections::HashMap;
fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("๐ง Advanced Quantization Example");
let device = Device::cuda_if_available(0).unwrap_or(Device::Cpu);
println!("๐ฑ Using device: {:?}", device);
let config = QuantizationConfig {
quantization_type: QuantizationType::Int8,
calibration_samples: 256,
calibration_method: "kl_divergence".to_string(), percentile: 99.5,
symmetric: true,
layer_config: HashMap::new(),
skip_layers: vec![
"embedding".to_string(),
"norm".to_string(),
"bias".to_string(),
],
quantize_bias: false,
block_wise: true, block_size: 512, advanced_stats: true, entropy_bins: 4096, kl_threshold: 0.05, };
println!("โ๏ธ Quantization Configuration:");
println!(" โข Type: {:?}", config.quantization_type);
println!(
" โข Calibration: {} with {} samples",
config.calibration_method, config.calibration_samples
);
println!(
" โข Block-wise: {} (block size: {})",
config.block_wise, config.block_size
);
println!(" โข Advanced stats: {}", config.advanced_stats);
let engine = QuantizationEngine::new(config, device.clone());
let model_path = "./model"; let load_options = LoadOptions {
device: device.clone(),
dtype: DType::F16,
use_mmap: true,
validate_cuda: false,
progress: Some(Box::new(|event| {
if let mlmf::progress::ProgressEvent::Status { message } = event {
println!("๐ Loading: {}", message);
}
})),
smart_mapping_oracle: None,
};
println!("\n๐ Loading model...");
let model = match load_model(model_path, load_options) {
Ok(model) => model,
Err(e) => {
println!("โ Failed to load model: {}", e);
println!(
"๐ก This example requires a valid model file at: {}",
model_path
);
return Ok(()); }
};
println!("โ
Model loaded with {} tensors", model.raw_tensors.len());
let mut context = QuantizationContext::new();
context.add_metadata(
"quantization_version".to_string(),
serde_json::json!("2.0-advanced"),
);
context.add_metadata(
"optimization_target".to_string(),
serde_json::json!("inference_speed"),
);
for tensor_name in model.raw_tensors.keys() {
if tensor_name.contains("attention") {
let attention_scheme = QuantizationScheme {
quant_type: QuantizationType::Int8,
scale: 0.01, zero_point: 0,
symmetric: true,
range: (-1.0, 1.0),
};
context.set_layer_scheme(tensor_name.clone(), attention_scheme);
println!(
"๐ฏ Set precise quantization for attention layer: {}",
tensor_name
);
}
}
println!("\n๐ Starting advanced quantization...");
let progress_callback = Some(Box::new(|event| {
if let mlmf::progress::ProgressEvent::Status { message } = event {
println!("โก {}", message);
}
})
as Box<dyn Fn(&mlmf::progress::ProgressEvent) + Send + Sync>);
let quantized_tensors =
engine.quantize_with_context(&model, &mut context, progress_callback)?;
println!("\n๐ Quantization Results:");
println!(" โข Tensors quantized: {}", quantized_tensors.len());
println!(
" โข Overall compression ratio: {:.2}x",
context.metrics.compression_ratio
);
println!(
" โข Quantization time: {:.2}s",
context.metrics.quantization_time
);
println!(" โข Average error: {:.4}", context.metrics.avg_error);
println!("\n๐ Per-Layer Compression Ratios:");
let mut sorted_layers: Vec<_> = context.metrics.layer_compression.iter().collect();
sorted_layers.sort_by(|a, b| b.1.partial_cmp(a.1).unwrap_or(std::cmp::Ordering::Equal));
for (layer_name, ratio) in sorted_layers.iter().take(5) {
println!(" โข {}: {:.2}x", layer_name, ratio);
}
println!("\n๐ Advanced Tensor Statistics (sample):");
for (tensor_name, stats) in context.tensor_stats.iter().take(3) {
println!(" ๐ {}:", tensor_name);
println!(" Range: [{:.4}, {:.4}]", stats.min_val, stats.max_val);
println!(
" Mean ยฑ Std: {:.4} ยฑ {:.4}",
stats.mean_val, stats.std_val
);
if let Some(ref percentiles) = stats.percentiles {
println!(
" Percentiles [P1, P5, P95, P99]: {:?}",
percentiles
.iter()
.map(|p| format!("{:.4}", p))
.collect::<Vec<_>>()
);
}
if let Some(ref kl_scores) = stats.kl_scores {
println!(" KL Divergence scores: {:?}", kl_scores);
}
}
println!("\n๐ Quantization Metadata:");
for (key, value) in &context.metadata {
println!(" โข {}: {}", key, value);
}
println!("\n๐งช Comparing Calibration Methods:");
demonstrate_calibration_methods(&engine, &model)?;
println!("\nโ
Advanced quantization demonstration completed!");
Ok(())
}
fn demonstrate_calibration_methods(
base_engine: &QuantizationEngine,
model: &mlmf::LoadedModel,
) -> Result<(), Box<dyn std::error::Error>> {
let methods = vec!["minmax", "percentile", "entropy", "kl_divergence"];
for method in methods {
let mut config = base_engine.config.clone();
config.calibration_method = method.to_string();
config.advanced_stats = false; config.block_wise = false;
let engine = QuantizationEngine::new(config, base_engine.device.clone());
let mut context = QuantizationContext::new();
let start_time = std::time::Instant::now();
let quantized_tensors = engine.quantize_with_context(model, &mut context, None)?;
let duration = start_time.elapsed().as_secs_f64();
println!(
" ๐ {}: {:.2}x compression, {:.2}s",
method, context.metrics.compression_ratio, duration
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_advanced_config_creation() {
let config = QuantizationConfig {
quantization_type: QuantizationType::Int8,
calibration_method: "kl_divergence".to_string(),
block_wise: true,
block_size: 1024,
advanced_stats: true,
entropy_bins: 2048,
kl_threshold: 0.1,
..Default::default()
};
assert_eq!(config.calibration_method, "kl_divergence");
assert!(config.block_wise);
assert!(config.advanced_stats);
}
#[test]
fn test_quantization_context() {
let mut context = QuantizationContext::new();
context.add_metadata("test_key".to_string(), serde_json::json!("test_value"));
assert_eq!(
context.get_metadata("test_key"),
Some(&serde_json::json!("test_value"))
);
}
}