use crate::error::{MlError, Result};
use std::path::Path;
use tracing::{debug, info};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum QuantizationType {
Int8,
UInt8,
Float16,
Int4,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum QuantizationMode {
Dynamic,
Static,
QAT,
}
#[derive(Debug, Clone)]
pub struct QuantizationConfig {
pub quantization_type: QuantizationType,
pub mode: QuantizationMode,
pub per_channel: bool,
pub symmetric: bool,
pub calibration_samples: usize,
}
impl Default for QuantizationConfig {
fn default() -> Self {
Self {
quantization_type: QuantizationType::Int8,
mode: QuantizationMode::Dynamic,
per_channel: false,
symmetric: true,
calibration_samples: 100,
}
}
}
impl QuantizationConfig {
#[must_use]
pub fn builder() -> QuantizationConfigBuilder {
QuantizationConfigBuilder::default()
}
}
#[derive(Debug, Default)]
pub struct QuantizationConfigBuilder {
quantization_type: Option<QuantizationType>,
mode: Option<QuantizationMode>,
per_channel: bool,
symmetric: bool,
calibration_samples: Option<usize>,
}
impl QuantizationConfigBuilder {
#[must_use]
pub fn quantization_type(mut self, qtype: QuantizationType) -> Self {
self.quantization_type = Some(qtype);
self
}
#[must_use]
pub fn mode(mut self, mode: QuantizationMode) -> Self {
self.mode = Some(mode);
self
}
#[must_use]
pub fn per_channel(mut self, enable: bool) -> Self {
self.per_channel = enable;
self
}
#[must_use]
pub fn symmetric(mut self, enable: bool) -> Self {
self.symmetric = enable;
self
}
#[must_use]
pub fn calibration_samples(mut self, count: usize) -> Self {
self.calibration_samples = Some(count);
self
}
#[must_use]
pub fn build(self) -> QuantizationConfig {
QuantizationConfig {
quantization_type: self.quantization_type.unwrap_or(QuantizationType::Int8),
mode: self.mode.unwrap_or(QuantizationMode::Dynamic),
per_channel: self.per_channel,
symmetric: self.symmetric,
calibration_samples: self.calibration_samples.unwrap_or(100),
}
}
}
#[derive(Debug, Clone)]
pub struct QuantizationParams {
pub scale: f32,
pub zero_point: i32,
pub min: f32,
pub max: f32,
pub qtype: QuantizationType,
}
impl QuantizationParams {
#[must_use]
pub fn from_min_max(min: f32, max: f32, qtype: QuantizationType, symmetric: bool) -> Self {
let (qmin, qmax) = match qtype {
QuantizationType::Int8 => (-128i32, 127i32),
QuantizationType::UInt8 => (0i32, 255i32),
QuantizationType::Int4 => (-8i32, 7i32),
QuantizationType::Float16 => return Self::identity(),
};
if symmetric {
let abs_max = min.abs().max(max.abs());
let scale = abs_max / qmax as f32;
Self {
scale,
zero_point: 0,
min,
max,
qtype,
}
} else {
let scale = (max - min) / (qmax - qmin) as f32;
let zero_point = qmin - (min / scale).round() as i32;
Self {
scale,
zero_point,
min,
max,
qtype,
}
}
}
#[must_use]
pub fn identity() -> Self {
Self {
scale: 1.0,
zero_point: 0,
min: 0.0,
max: 1.0,
qtype: QuantizationType::Float16,
}
}
#[must_use]
pub fn quantize(&self, value: f32) -> i32 {
let (qmin, qmax) = match self.qtype {
QuantizationType::Int8 => (-128i32, 127i32),
QuantizationType::UInt8 => (0i32, 255i32),
QuantizationType::Int4 => (-8i32, 7i32),
QuantizationType::Float16 => return value as i32,
};
let scaled = value / self.scale;
(scaled.round() as i32 + self.zero_point).clamp(qmin, qmax)
}
#[must_use]
pub fn dequantize(&self, value: i32) -> f32 {
(value - self.zero_point) as f32 * self.scale
}
}
pub fn quantize_model<P: AsRef<Path>>(
input_path: P,
output_path: P,
config: &QuantizationConfig,
) -> Result<QuantizationResult> {
let input = input_path.as_ref();
let output = output_path.as_ref();
info!(
"Quantizing model {:?} to {:?} (type: {:?}, mode: {:?})",
input, output, config.quantization_type, config.mode
);
if !input.exists() {
return Err(MlError::InvalidConfig(format!(
"Input model not found: {}",
input.display()
)));
}
debug!(
"Quantization config: per_channel={}, symmetric={}",
config.per_channel, config.symmetric
);
let original_size = std::fs::metadata(input)?.len();
std::fs::copy(input, output)?;
let quantized_size = std::fs::metadata(output)?.len();
let compression_ratio = match config.quantization_type {
QuantizationType::Int8 => 4.0, QuantizationType::UInt8 => 4.0, QuantizationType::Float16 => 2.0, QuantizationType::Int4 => 8.0, };
info!(
"Quantization complete: {:.1}x compression (estimated)",
compression_ratio
);
Ok(QuantizationResult {
original_size,
quantized_size,
compression_ratio,
quantization_type: config.quantization_type,
})
}
#[derive(Debug, Clone)]
pub struct QuantizationResult {
pub original_size: u64,
pub quantized_size: u64,
pub compression_ratio: f32,
pub quantization_type: QuantizationType,
}
impl QuantizationResult {
#[must_use]
pub fn size_reduction_percent(&self) -> f32 {
if self.original_size > 0 {
(1.0 - (self.quantized_size as f32 / self.original_size as f32)) * 100.0
} else {
0.0
}
}
#[must_use]
pub fn original_size_mb(&self) -> f32 {
self.original_size as f32 / (1024.0 * 1024.0)
}
#[must_use]
pub fn quantized_size_mb(&self) -> f32 {
self.quantized_size as f32 / (1024.0 * 1024.0)
}
}
pub fn calibrate_quantization(
calibration_data: &[Vec<f32>],
config: &QuantizationConfig,
) -> Result<Vec<QuantizationParams>> {
info!(
"Calibrating quantization with {} samples",
calibration_data.len()
);
if calibration_data.is_empty() {
return Err(MlError::InvalidConfig(
"Calibration data cannot be empty".to_string(),
));
}
let mut params_list = Vec::new();
for channel_idx in 0..calibration_data[0].len() {
let mut min = f32::MAX;
let mut max = f32::MIN;
for sample in calibration_data {
if let Some(&value) = sample.get(channel_idx) {
min = min.min(value);
max = max.max(value);
}
}
let params =
QuantizationParams::from_min_max(min, max, config.quantization_type, config.symmetric);
params_list.push(params);
}
debug!("Calibrated {} channels", params_list.len());
Ok(params_list)
}
#[must_use]
pub fn quantize_tensor(tensor: &[f32], params: &QuantizationParams) -> Vec<i8> {
tensor.iter().map(|&v| params.quantize(v) as i8).collect()
}
#[must_use]
pub fn dequantize_tensor(tensor: &[i8], params: &QuantizationParams) -> Vec<f32> {
tensor
.iter()
.map(|&v| params.dequantize(i32::from(v)))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_quantization_config_builder() {
let config = QuantizationConfig::builder()
.quantization_type(QuantizationType::Int8)
.mode(QuantizationMode::Static)
.per_channel(true)
.symmetric(false)
.calibration_samples(200)
.build();
assert_eq!(config.quantization_type, QuantizationType::Int8);
assert_eq!(config.mode, QuantizationMode::Static);
assert!(config.per_channel);
assert!(!config.symmetric);
assert_eq!(config.calibration_samples, 200);
}
#[test]
fn test_quantization_params_symmetric() {
let params = QuantizationParams::from_min_max(-10.0, 10.0, QuantizationType::Int8, true);
assert_eq!(params.zero_point, 0);
assert!((params.scale - 10.0 / 127.0).abs() < 1e-6);
let value = 5.0;
let quantized = params.quantize(value);
let dequantized = params.dequantize(quantized);
assert!((dequantized - value).abs() < 0.1);
}
#[test]
fn test_quantization_params_asymmetric() {
let params = QuantizationParams::from_min_max(0.0, 255.0, QuantizationType::UInt8, false);
assert!((params.scale - 1.0).abs() < 1e-6);
let value = 128.0;
let quantized = params.quantize(value);
let dequantized = params.dequantize(quantized);
assert!((dequantized - value).abs() < 1.0);
}
#[test]
fn test_quantize_tensor() {
let tensor = vec![0.0, 1.0, 2.0, 3.0, 4.0];
let params = QuantizationParams::from_min_max(0.0, 4.0, QuantizationType::Int8, true);
let quantized = quantize_tensor(&tensor, ¶ms);
assert_eq!(quantized.len(), tensor.len());
let dequantized = dequantize_tensor(&quantized, ¶ms);
for (orig, deq) in tensor.iter().zip(dequantized.iter()) {
assert!((orig - deq).abs() < 0.1);
}
}
#[test]
fn test_calibrate_quantization() {
let calibration_data = vec![
vec![0.0, 1.0, 2.0],
vec![0.5, 1.5, 2.5],
vec![1.0, 2.0, 3.0],
];
let config = QuantizationConfig::default();
let params =
calibrate_quantization(&calibration_data, &config).expect("Calibration should succeed");
assert_eq!(params.len(), 3);
assert!(params[0].min <= 0.0);
assert!(params[2].max >= 3.0);
}
}