use crate::{NeuralResult, SklearsError};
use scirs2_core::ndarray::Array2;
use sklears_core::types::{Float, FloatBounds};
use std::collections::HashMap;
use std::marker::PhantomData;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum QuantizationType {
#[default]
INT8,
INT16,
INT4,
Binary,
Ternary,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum QuantizationStrategy {
#[default]
PostTraining,
QuantizationAware,
Dynamic,
Static,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum Granularity {
#[default]
PerTensor,
PerChannel,
PerGroup {
group_size: usize,
},
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum CalibrationMethod {
#[default]
MinMax,
Percentile {
percentile: f64,
},
Entropy,
MSE,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct QuantizationConfig {
pub quantization_type: QuantizationType,
pub strategy: QuantizationStrategy,
pub granularity: Granularity,
pub calibration_method: CalibrationMethod,
pub symmetric: bool,
pub quantize_weights: bool,
pub quantize_activations: bool,
pub skip_layers: Vec<String>,
pub qat_epochs: usize,
pub qat_learning_rate: Float,
pub fake_quantize: bool,
pub observer_momentum: Float,
pub calibration_samples: usize,
pub random_state: Option<u64>,
}
impl Default for QuantizationConfig {
fn default() -> Self {
Self {
quantization_type: QuantizationType::INT8,
strategy: QuantizationStrategy::PostTraining,
granularity: Granularity::PerTensor,
calibration_method: CalibrationMethod::MinMax,
symmetric: true,
quantize_weights: true,
quantize_activations: true,
skip_layers: vec!["input".to_string(), "output".to_string()],
qat_epochs: 10,
qat_learning_rate: 0.0001,
fake_quantize: false,
observer_momentum: 0.1,
calibration_samples: 1000,
random_state: None,
}
}
}
impl QuantizationConfig {
pub fn quantization_type(mut self, qtype: QuantizationType) -> Self {
self.quantization_type = qtype;
self
}
pub fn strategy(mut self, strategy: QuantizationStrategy) -> Self {
self.strategy = strategy;
self
}
pub fn granularity(mut self, granularity: Granularity) -> Self {
self.granularity = granularity;
self
}
pub fn calibration_method(mut self, method: CalibrationMethod) -> Self {
self.calibration_method = method;
self
}
pub fn symmetric(mut self, symmetric: bool) -> Self {
self.symmetric = symmetric;
self
}
pub fn quantize_components(mut self, weights: bool, activations: bool) -> Self {
self.quantize_weights = weights;
self.quantize_activations = activations;
self
}
pub fn skip_layers(mut self, layers: Vec<String>) -> Self {
self.skip_layers = layers;
self
}
pub fn qat_config(mut self, epochs: usize, lr: Float) -> Self {
self.strategy = QuantizationStrategy::QuantizationAware;
self.qat_epochs = epochs;
self.qat_learning_rate = lr;
self
}
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct QuantizationParams {
pub scale: f64,
pub zero_point: i32,
pub min_val: f64,
pub max_val: f64,
pub n_levels: u32,
}
impl QuantizationParams {
pub fn new(scale: f64, zero_point: i32, min_val: f64, max_val: f64, n_levels: u32) -> Self {
Self {
scale,
zero_point,
min_val,
max_val,
n_levels,
}
}
pub fn symmetric(max_abs: f64, n_levels: u32) -> Self {
let scale = 2.0 * max_abs / (n_levels - 1) as f64;
Self {
scale,
zero_point: 0,
min_val: -max_abs,
max_val: max_abs,
n_levels,
}
}
pub fn asymmetric(min_val: f64, max_val: f64, n_levels: u32) -> Self {
let scale = (max_val - min_val) / (n_levels - 1) as f64;
let zero_point = (-min_val / scale).round() as i32;
Self {
scale,
zero_point,
min_val,
max_val,
n_levels,
}
}
}
#[derive(Debug, Clone)]
pub struct QuantizedTensor {
pub data: Array2<i32>,
pub params: QuantizationParams,
pub original_shape: Vec<usize>,
}
impl QuantizedTensor {
pub fn new(data: Array2<i32>, params: QuantizationParams, original_shape: Vec<usize>) -> Self {
Self {
data,
params,
original_shape,
}
}
pub fn dequantize<T: FloatBounds>(&self) -> NeuralResult<Array2<T>> {
let mut result = Array2::zeros(self.data.dim());
for (i, &quantized_val) in self.data.iter().enumerate() {
let dequantized = self.params.scale * (quantized_val - self.params.zero_point) as f64;
result.as_slice_mut().expect("non-contiguous array")[i] =
T::from(dequantized).unwrap_or_else(|| T::zero());
}
Ok(result)
}
pub fn compression_ratio(&self) -> f64 {
32.0 / 8.0
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)] pub struct Observer {
min_val: f64,
max_val: f64,
count: usize,
momentum: f64,
histogram: Vec<u64>,
histogram_bins: usize,
}
impl Observer {
pub fn new(momentum: f64, histogram_bins: usize) -> Self {
Self {
min_val: f64::INFINITY,
max_val: f64::NEG_INFINITY,
count: 0,
momentum,
histogram: vec![0; histogram_bins],
histogram_bins,
}
}
pub fn update<T: FloatBounds>(&mut self, data: &Array2<T>) {
let is_first_update = self.count == 0;
for &val in data.iter() {
let val_f64 = val.to_f64().unwrap_or(0.0);
if is_first_update && self.min_val == f64::INFINITY {
self.min_val = val_f64;
self.max_val = val_f64;
} else {
self.min_val = self.min_val.min(val_f64);
self.max_val = self.max_val.max(val_f64);
}
self.update_histogram(val_f64);
}
self.count += 1;
}
fn update_histogram(&mut self, val: f64) {
if self.min_val < self.max_val {
let range = self.max_val - self.min_val;
let normalized = (val - self.min_val) / range;
let bin_idx = ((normalized * (self.histogram_bins - 1) as f64).floor() as usize)
.min(self.histogram_bins - 1);
self.histogram[bin_idx] += 1;
}
}
pub fn get_quantization_params(
&self,
method: &CalibrationMethod,
symmetric: bool,
n_levels: u32,
) -> QuantizationParams {
match method {
CalibrationMethod::MinMax => {
if symmetric {
let max_abs = self.min_val.abs().max(self.max_val.abs());
QuantizationParams::symmetric(max_abs, n_levels)
} else {
QuantizationParams::asymmetric(self.min_val, self.max_val, n_levels)
}
}
CalibrationMethod::Percentile { percentile } => {
let (min_p, max_p) = self.compute_percentile(*percentile);
if symmetric {
let max_abs = min_p.abs().max(max_p.abs());
QuantizationParams::symmetric(max_abs, n_levels)
} else {
QuantizationParams::asymmetric(min_p, max_p, n_levels)
}
}
CalibrationMethod::Entropy => {
let (min_opt, max_opt) = self.compute_optimal_clipping(n_levels);
QuantizationParams::asymmetric(min_opt, max_opt, n_levels)
}
CalibrationMethod::MSE => {
let (min_mse, max_mse) = self.compute_mse_optimal();
QuantizationParams::asymmetric(min_mse, max_mse, n_levels)
}
}
}
fn compute_percentile(&self, percentile: f64) -> (f64, f64) {
let total_count: u64 = self.histogram.iter().sum();
if total_count == 0 {
return (self.min_val, self.max_val);
}
let target_low = ((100.0 - percentile) / 2.0 / 100.0 * total_count as f64) as u64;
let target_high = (((100.0 + percentile) / 2.0 / 100.0) * total_count as f64) as u64;
let mut cumsum = 0;
let mut min_p = self.min_val;
let mut max_p = self.max_val;
let bin_width = (self.max_val - self.min_val) / self.histogram_bins as f64;
for (i, &count) in self.histogram.iter().enumerate() {
cumsum += count;
if cumsum >= target_low && min_p == self.min_val {
min_p = self.min_val + i as f64 * bin_width;
}
if cumsum >= target_high {
max_p = self.min_val + (i + 1) as f64 * bin_width;
break;
}
}
(min_p, max_p)
}
fn compute_optimal_clipping(&self, _n_levels: u32) -> (f64, f64) {
let (min_p, max_p) = self.compute_percentile(99.9);
(min_p, max_p)
}
fn compute_mse_optimal(&self) -> (f64, f64) {
(self.min_val, self.max_val)
}
}
#[allow(dead_code)] pub struct Quantizer<T: FloatBounds> {
config: QuantizationConfig,
layer_params: HashMap<String, QuantizationParams>,
observers: HashMap<String, Observer>,
calibration_data: HashMap<String, Vec<Array2<T>>>,
_phantom: PhantomData<T>,
}
impl<T: FloatBounds> Quantizer<T> {
pub fn new(config: QuantizationConfig) -> Self {
Self {
config,
layer_params: HashMap::new(),
observers: HashMap::new(),
calibration_data: HashMap::new(),
_phantom: PhantomData,
}
}
pub fn calibrate(&mut self, model_data: &HashMap<String, Array2<T>>) -> NeuralResult<()> {
for layer_name in model_data.keys() {
if !self.config.skip_layers.contains(layer_name) {
let observer = Observer::new(self.config.observer_momentum, 256);
self.observers.insert(layer_name.clone(), observer);
}
}
for (layer_name, data) in model_data {
if let Some(observer) = self.observers.get_mut(layer_name) {
observer.update(data);
}
}
for (layer_name, observer) in &self.observers {
let n_levels = self.get_quantization_levels();
let params = observer.get_quantization_params(
&self.config.calibration_method,
self.config.symmetric,
n_levels,
);
self.layer_params.insert(layer_name.clone(), params);
}
Ok(())
}
pub fn quantize_tensor(
&self,
tensor: &Array2<T>,
layer_name: &str,
) -> NeuralResult<QuantizedTensor> {
let params =
self.layer_params
.get(layer_name)
.ok_or_else(|| SklearsError::InvalidParameter {
name: "layer_name".to_string(),
reason: format!("No quantization parameters found for layer {}", layer_name),
})?;
let quantized_data = self.apply_quantization(tensor, params)?;
let original_shape = tensor.shape().to_vec();
Ok(QuantizedTensor::new(
quantized_data,
params.clone(),
original_shape,
))
}
fn apply_quantization(
&self,
tensor: &Array2<T>,
params: &QuantizationParams,
) -> NeuralResult<Array2<i32>> {
let mut quantized = Array2::zeros(tensor.dim());
for (i, &val) in tensor.iter().enumerate() {
let val_f64 = val.to_f64().unwrap_or(0.0);
let quantized_val = (val_f64 / params.scale).round() as i32 + params.zero_point;
let clamped = quantized_val
.max(-(params.n_levels as i32 / 2))
.min(params.n_levels as i32 / 2 - 1);
quantized.as_slice_mut().expect("non-contiguous array")[i] = clamped;
}
Ok(quantized)
}
pub fn dequantize_tensor(&self, quantized: &QuantizedTensor) -> NeuralResult<Array2<T>> {
quantized.dequantize()
}
pub fn quantize_weights(
&mut self,
weights: &HashMap<String, Array2<T>>,
) -> NeuralResult<HashMap<String, QuantizedTensor>> {
if !self.config.quantize_weights {
return Err(SklearsError::InvalidParameter {
name: "quantize_weights".to_string(),
reason: "Weight quantization is disabled".to_string(),
});
}
if self.layer_params.is_empty() {
self.calibrate(weights)?;
}
let mut quantized_weights = HashMap::new();
for (layer_name, weight) in weights {
if !self.config.skip_layers.contains(layer_name) {
let quantized = self.quantize_tensor(weight, layer_name)?;
quantized_weights.insert(layer_name.clone(), quantized);
}
}
Ok(quantized_weights)
}
pub fn fake_quantize_tensor(
&self,
tensor: &Array2<T>,
layer_name: &str,
) -> NeuralResult<Array2<T>> {
if !self.config.fake_quantize {
return Ok(tensor.clone());
}
let quantized = self.quantize_tensor(tensor, layer_name)?;
self.dequantize_tensor(&quantized)
}
fn get_quantization_levels(&self) -> u32 {
match self.config.quantization_type {
QuantizationType::INT8 => 256,
QuantizationType::INT16 => 65536,
QuantizationType::INT4 => 16,
QuantizationType::Binary => 2,
QuantizationType::Ternary => 3,
}
}
pub fn compute_quantization_error(
&self,
original: &Array2<T>,
quantized: &QuantizedTensor,
) -> NeuralResult<QuantizationMetrics> {
let reconstructed = self.dequantize_tensor(quantized)?;
let diff = original - &reconstructed;
let mse = diff
.mapv(|x| x * x)
.mean()
.expect("mean should not fail on non-empty array")
.to_f64()
.unwrap_or(0.0);
let signal_power = original
.mapv(|x| x * x)
.mean()
.expect("value should be present")
.to_f64()
.unwrap_or(0.0);
let snr_db = if mse > 0.0 {
10.0 * (signal_power / mse).log10()
} else {
f64::INFINITY
};
let max_val = original
.iter()
.map(|&x| x.abs())
.fold(T::from(0.0).unwrap_or_else(|| T::zero()), |a, b| a.max(b));
let max_val_f64 = max_val.to_f64().unwrap_or(0.0);
let psnr_db = if mse > 0.0 {
20.0 * (max_val_f64 / mse.sqrt()).log10()
} else {
f64::INFINITY
};
let compression_ratio = quantized.compression_ratio();
Ok(QuantizationMetrics {
mse,
snr_db,
psnr_db,
compression_ratio,
original_size: original.len() * std::mem::size_of::<T>(),
quantized_size: quantized.data.len() * std::mem::size_of::<i32>(),
})
}
pub fn analyze_layer_sensitivity(
&mut self,
layers_data: &HashMap<String, Array2<T>>,
) -> NeuralResult<HashMap<String, f64>> {
let mut sensitivity_scores = HashMap::new();
for (layer_name, data) in layers_data {
if self.config.skip_layers.contains(layer_name) {
continue;
}
let quantized = self.quantize_tensor(data, layer_name)?;
let metrics = self.compute_quantization_error(data, &quantized)?;
sensitivity_scores.insert(layer_name.clone(), metrics.mse);
}
Ok(sensitivity_scores)
}
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct QuantizationMetrics {
pub mse: f64,
pub snr_db: f64,
pub psnr_db: f64,
pub compression_ratio: f64,
pub original_size: usize,
pub quantized_size: usize,
}
impl QuantizationMetrics {
pub fn is_acceptable(&self, min_snr_db: f64, max_mse: f64) -> bool {
self.snr_db >= min_snr_db && self.mse <= max_mse
}
pub fn memory_savings(&self) -> f64 {
1.0 - (self.quantized_size as f64 / self.original_size as f64)
}
}
pub mod utils {
use super::*;
pub fn int8_ptq_config() -> QuantizationConfig {
QuantizationConfig::default()
.quantization_type(QuantizationType::INT8)
.strategy(QuantizationStrategy::PostTraining)
.symmetric(true)
}
pub fn int8_qat_config(epochs: usize, lr: f64) -> QuantizationConfig {
QuantizationConfig::default()
.quantization_type(QuantizationType::INT8)
.qat_config(epochs, lr)
}
pub fn dynamic_quantization_config() -> QuantizationConfig {
QuantizationConfig::default()
.strategy(QuantizationStrategy::Dynamic)
.quantize_components(true, false) }
pub fn per_channel_config() -> QuantizationConfig {
QuantizationConfig::default()
.granularity(Granularity::PerChannel)
.symmetric(false)
}
pub fn binary_quantization_config() -> QuantizationConfig {
QuantizationConfig::default()
.quantization_type(QuantizationType::Binary)
.symmetric(true)
}
pub fn evaluate_accuracy_impact<T: FloatBounds>(
original_accuracy: f64,
quantized_accuracy: f64,
) -> f64 {
(original_accuracy - quantized_accuracy) / original_accuracy
}
pub fn select_optimal_parameters(
sensitivity_scores: &HashMap<String, f64>,
target_compression: f64,
) -> Vec<String> {
let mut layers: Vec<(String, f64)> = sensitivity_scores
.iter()
.map(|(name, &score)| (name.clone(), score))
.collect();
layers.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
let n_layers_to_quantize = (layers.len() as f64 * target_compression) as usize;
layers
.into_iter()
.take(n_layers_to_quantize)
.map(|(name, _)| name)
.collect()
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use approx;
use scirs2_core::ndarray::arr2;
#[test]
fn test_quantization_config() {
let config = QuantizationConfig::default()
.quantization_type(QuantizationType::INT8)
.strategy(QuantizationStrategy::PostTraining)
.symmetric(false);
assert_eq!(config.quantization_type, QuantizationType::INT8);
assert_eq!(config.strategy, QuantizationStrategy::PostTraining);
assert!(!config.symmetric);
}
#[test]
fn test_quantization_params() {
let params = QuantizationParams::symmetric(1.0, 256);
assert_eq!(params.zero_point, 0);
approx::assert_abs_diff_eq!(params.scale, 2.0 / 255.0, epsilon = 1e-10);
let params = QuantizationParams::asymmetric(-1.0, 2.0, 256);
approx::assert_abs_diff_eq!(params.scale, 3.0 / 255.0, epsilon = 1e-10);
assert!(params.zero_point > 0);
}
#[test]
fn test_observer() {
let mut observer = Observer::new(0.1, 100);
let data = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
observer.update(&data);
assert_eq!(observer.min_val, 1.0);
assert_eq!(observer.max_val, 4.0);
let params = observer.get_quantization_params(&CalibrationMethod::MinMax, true, 256);
assert_eq!(params.zero_point, 0); }
#[test]
fn test_quantization_dequantization() {
let config = QuantizationConfig::default();
let mut quantizer = Quantizer::<f64>::new(config);
let data = arr2(&[[1.0, -1.0], [0.5, -0.5]]);
let mut layer_data = HashMap::new();
layer_data.insert("test_layer".to_string(), data.clone());
quantizer
.calibrate(&layer_data)
.expect("operation should succeed");
let quantized = quantizer
.quantize_tensor(&data, "test_layer")
.expect("operation should succeed");
let reconstructed = quantizer
.dequantize_tensor(&quantized)
.expect("operation should succeed");
for (orig, recon) in data.iter().zip(reconstructed.iter()) {
approx::assert_abs_diff_eq!(orig, recon, epsilon = 0.1);
}
}
#[test]
fn test_quantization_metrics() {
let original = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
let _reconstructed = arr2(&[[1.1, 1.9], [3.1, 3.9]]);
let config = QuantizationConfig::default();
let quantizer = Quantizer::<f64>::new(config);
let quantized_data = arr2(&[[110, 190], [310, 390]]);
let params = QuantizationParams::symmetric(4.0, 256);
let quantized = QuantizedTensor::new(quantized_data, params, vec![2, 2]);
let metrics = quantizer
.compute_quantization_error(&original, &quantized)
.expect("operation should succeed");
assert!(metrics.mse >= 0.0);
assert!(metrics.compression_ratio > 1.0);
}
#[test]
fn test_utility_functions() {
let config = utils::int8_ptq_config();
assert_eq!(config.quantization_type, QuantizationType::INT8);
assert_eq!(config.strategy, QuantizationStrategy::PostTraining);
let qat_config = utils::int8_qat_config(10, 0.001);
assert_eq!(qat_config.strategy, QuantizationStrategy::QuantizationAware);
assert_eq!(qat_config.qat_epochs, 10);
let accuracy_impact = utils::evaluate_accuracy_impact::<f64>(0.95, 0.92);
approx::assert_abs_diff_eq!(accuracy_impact, 0.0315789, epsilon = 1e-6);
}
#[test]
fn test_sensitivity_analysis() {
let mut sensitivity_scores = HashMap::new();
sensitivity_scores.insert("layer1".to_string(), 0.1);
sensitivity_scores.insert("layer2".to_string(), 0.05);
sensitivity_scores.insert("layer3".to_string(), 0.2);
let selected = utils::select_optimal_parameters(&sensitivity_scores, 0.67);
assert_eq!(selected.len(), 2); assert!(selected.contains(&"layer2".to_string())); }
#[test]
fn test_quantized_tensor() {
let data = arr2(&[[100, 150], [200, 250]]);
let params = QuantizationParams::symmetric(2.0, 256);
let quantized = QuantizedTensor::new(data, params, vec![2, 2]);
let dequantized = quantized
.dequantize::<f64>()
.expect("operation should succeed");
assert!(dequantized.dim() == (2, 2));
let ratio = quantized.compression_ratio();
assert_eq!(ratio, 4.0); }
}