Files
foxhunt/ml/src/memory_optimization/qat.rs
jgrusewski 73b9ca0659 fix(clippy): Fix 17 critical float_arithmetic warnings in load_tests
- Added safe_div(), safe_mul(), and safe_add() helper functions
- All helpers check for NaN, infinity, and division by zero
- Replaced direct float operations with safe wrappers
- Fixed percentile calculations (lines 86-89)
- Fixed success rate calculation (line 101)
- Fixed throughput calculation (line 107)
- Fixed all latency metric conversions (lines 133-154)
- Fixed P99 latency display (lines 177, 182)
- Fixed order quantity/price calculations (lines 215-216)

All 17 float_arithmetic warnings in lib.rs now resolved.
Part 1/2: 9 warnings requested, 17 actually fixed.
2025-10-23 11:54:56 +02:00

1360 lines
46 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Quantization-Aware Training (QAT) Infrastructure for TFT
//!
//! Implements fake quantization layers and observers to simulate INT8 quantization
//! during training, allowing the model to adapt to quantization noise and maintain
//! accuracy when deployed with INT8 weights.
//!
//! # Key Components
//! - `QATConfig`: Configuration for quantization-aware training
//! - `FakeQuantize`: Layer that simulates quantization during training
//! - `QuantizationObserver`: Trait for tracking min/max statistics
//! - `MinMaxObserver`: Observer implementation with EMA smoothing
//!
//! # Usage Example
//! ```ignore
//! use ml::memory_optimization::qat::{QATConfig, FakeQuantize, MinMaxObserver};
//! use candle_core::{Device, Tensor};
//!
//! let config = QATConfig::default();
//! let device = Device::cuda_if_available(0)?;
//!
//! // Create fake quantization layer with observer
//! let mut fake_quant = FakeQuantize::new(
//! config.quant_type,
//! config.symmetric,
//! config.per_channel,
//! device.clone()
//! )?;
//!
//! // During training: forward pass quantizes then dequantizes (gradient flows through)
//! let input = Tensor::randn(0.0, 1.0, (32, 256), &device)?;
//! let output = fake_quant.forward(&input, true)?; // true = training mode
//!
//! // Observer tracks min/max statistics with EMA smoothing
//! println!("Observed range: [{:.3}, {:.3}]", fake_quant.min(), fake_quant.max());
//! ```
use candle_core::{Device, Tensor};
use serde::{Deserialize, Serialize};
use std::path::Path;
use std::sync::{Arc, Mutex};
use tracing::{debug, info};
use crate::memory_optimization::quantization::{QuantizationType, QuantizedTensor};
use crate::MLError;
/// Quantization-Aware Training Configuration
///
/// Controls how fake quantization is applied during training to simulate
/// INT8 deployment and adapt model weights to quantization noise.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QATConfig {
/// Quantization type (Int8, Int4, etc.)
pub quant_type: QuantizationType,
/// Symmetric vs asymmetric quantization
/// - Symmetric: Maps [-abs_max, abs_max] → [0, 255] with zero_point=127
/// - Asymmetric: Maps [min, max] → [0, 255] with learned zero_point
pub symmetric: bool,
/// Per-channel quantization (better accuracy, more memory)
/// - true: Separate scale/zero_point per output channel (Conv/Linear layers)
/// - false: Single scale/zero_point for entire tensor
pub per_channel: bool,
/// Number of calibration batches for observer initialization
/// Determines how many batches to collect statistics before starting fake quantization
pub calibration_batches: usize,
/// Enable fake quantization during forward pass
/// - true: Apply quantize→dequantize (training mode)
/// - false: Pass-through (validation/testing)
pub fake_quant_enabled: bool,
/// Observer update frequency (in batches)
/// Updates min/max statistics every N batches to reduce overhead
pub observer_update_frequency: usize,
/// EMA decay factor for running statistics (0.99 = slow adaptation, 0.9 = fast)
pub ema_decay: f32,
}
impl Default for QATConfig {
fn default() -> Self {
Self {
quant_type: QuantizationType::Int8,
symmetric: true,
per_channel: true,
calibration_batches: 100,
fake_quant_enabled: true,
observer_update_frequency: 10,
ema_decay: 0.99,
}
}
}
/// Observer for tracking activation statistics during calibration
///
/// Collects running min/max values using exponential moving average (EMA)
/// to determine optimal quantization parameters before QAT training.
#[derive(Debug, Clone)]
pub struct QuantizationObserver {
config: QATConfig,
device: Device,
/// Running minimum (EMA)
running_min: Arc<Mutex<Option<f32>>>,
/// Running maximum (EMA)
running_max: Arc<Mutex<Option<f32>>>,
/// Number of observations
num_observations: Arc<Mutex<usize>>,
/// Calibration complete flag
calibrated: Arc<Mutex<bool>>,
}
impl QuantizationObserver {
/// Create a new quantization observer
pub fn new(config: QATConfig, device: Device) -> Self {
Self {
config,
device,
running_min: Arc::new(Mutex::new(None)),
running_max: Arc::new(Mutex::new(None)),
num_observations: Arc::new(Mutex::new(0)),
calibrated: Arc::new(Mutex::new(false)),
}
}
/// Observe a batch of activations and update running statistics
///
/// Uses exponential moving average (EMA) to track min/max values:
/// - running_min = ema_decay * running_min + (1 - ema_decay) * batch_min
/// - running_max = ema_decay * running_max + (1 - ema_decay) * batch_max
///
/// # Arguments
/// * `activations` - Batch of activations to observe
///
/// # Returns
/// * `Ok(())` - Observation successful
/// * `Err(MLError)` - If tensor conversion fails
pub fn observe(&mut self, activations: &Tensor) -> Result<(), MLError> {
// Convert to F32 for statistics
let f32_activations = activations.to_dtype(DType::F32)?;
let flat = f32_activations.flatten_all()?;
let data = flat
.to_vec1::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to convert tensor to vec: {}", e)))?;
// Compute batch statistics
let batch_min = data.iter().cloned().fold(f32::INFINITY, f32::min);
let batch_max = data.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
// Update running statistics with EMA
let mut min_lock = self.running_min.lock().unwrap();
let mut max_lock = self.running_max.lock().unwrap();
let mut count_lock = self.num_observations.lock().unwrap();
match (*min_lock, *max_lock) {
(Some(current_min), Some(current_max)) => {
// EMA update: new = decay * old + (1 - decay) * new
let alpha = 1.0 - self.config.ema_decay;
*min_lock = Some(self.config.ema_decay * current_min + alpha * batch_min);
*max_lock = Some(self.config.ema_decay * current_max + alpha * batch_max);
}
_ => {
// First observation
*min_lock = Some(batch_min);
*max_lock = Some(batch_max);
}
}
*count_lock += 1;
// Mark as calibrated if we've seen enough batches
if *count_lock >= self.config.calibration_batches {
*self.calibrated.lock().unwrap() = true;
}
Ok(())
}
/// Check if calibration is complete
pub fn is_calibrated(&self) -> bool {
*self.calibrated.lock().unwrap()
}
/// Get calibrated min/max values
///
/// # Returns
/// * `Some((min, max))` - Calibrated min/max values
/// * `None` - Not calibrated yet
pub fn get_min_max(&self) -> Option<(f32, f32)> {
let min_lock = self.running_min.lock().unwrap();
let max_lock = self.running_max.lock().unwrap();
match (*min_lock, *max_lock) {
(Some(min), Some(max)) => Some((min, max)),
_ => None,
}
}
/// Get number of observations
pub fn num_observations(&self) -> usize {
*self.num_observations.lock().unwrap()
}
/// Reset observer statistics
pub fn reset(&mut self) {
*self.running_min.lock().unwrap() = None;
*self.running_max.lock().unwrap() = None;
*self.num_observations.lock().unwrap() = 0;
*self.calibrated.lock().unwrap() = false;
}
}
/// Fake quantization layer for QAT
///
/// Simulates INT8 quantization during forward pass while allowing gradients
/// to flow through during backward pass (Straight-Through Estimator).
///
/// # Forward Pass
/// 1. Quantize: x_q = clamp(round(x / scale + zero_point), 0, 255)
/// 2. Dequantize: x_dq = scale * (x_q - zero_point)
///
/// # Backward Pass
/// - Gradients flow through as if quantization didn't exist (STE)
/// - This allows the network to learn quantization-robust weights
pub struct FakeQuantize {
config: QATConfig,
device: Device,
/// Quantization scale
scale: f32,
/// Quantization zero point
zero_point: i8,
/// Min value (for calibration)
min_val: f32,
/// Max value (for calibration)
max_val: f32,
/// Training mode flag
training: bool,
}
impl FakeQuantize {
/// Create from calibrated observer
pub fn from_observer(observer: &QuantizationObserver) -> Result<Self, MLError> {
if !observer.is_calibrated() {
return Err(MLError::ModelError(
"Observer not calibrated. Run calibration phase first.".to_string(),
));
}
let (min_val, max_val) = observer.get_min_max().ok_or_else(|| {
MLError::ModelError("Observer has no min/max statistics".to_string())
})?;
// Calculate quantization parameters
let (scale, zero_point) = if observer.config.symmetric {
// Symmetric quantization: scale = max(abs(min), abs(max)) / 127
let abs_max = min_val.abs().max(max_val.abs());
let scale = abs_max / 127.0;
(scale, 127i8)
} else {
// Asymmetric quantization
let scale = (max_val - min_val) / 255.0;
let zero_point = (-min_val / scale).round() as i8;
(scale, zero_point)
};
Ok(Self {
config: observer.config.clone(),
device: observer.device.clone(),
scale,
zero_point,
min_val,
max_val,
training: true,
})
}
/// Create with explicit scale and zero point (for testing)
pub fn new(
config: QATConfig,
device: Device,
scale: f32,
zero_point: i8,
) -> Result<Self, MLError> {
Ok(Self {
config,
device,
scale,
zero_point,
min_val: 0.0,
max_val: 255.0 * scale,
training: true,
})
}
/// Forward pass with fake quantization
///
/// Simulates INT8 quantization during training:
/// 1. Quantize: x_q = clamp(round(x / scale + zero_point), 0, 255)
/// 2. Dequantize: x_dq = scale * (x_q - zero_point)
///
/// Gradients flow through using Straight-Through Estimator (STE).
///
/// # Arguments
/// * `input` - Input tensor (any dtype)
///
/// # Returns
/// * Fake-quantized tensor (same dtype as input)
pub fn forward(&self, input: &Tensor) -> Result<Tensor, MLError> {
if !self.training {
// Evaluation mode: no quantization simulation
return Ok(input.clone());
}
// Convert to F32 for quantization
let f32_input = input.to_dtype(DType::F32)?;
// Quantize: q = clamp(round((x / scale) + zero_point), 0, 255)
let input_device = f32_input.device();
let scale_tensor = Tensor::new(&[self.scale], input_device)?;
let zero_point_tensor = Tensor::new(&[self.zero_point as f32], input_device)?;
let scaled = f32_input.broadcast_div(&scale_tensor)?;
let shifted = scaled.broadcast_add(&zero_point_tensor)?;
let rounded = shifted
.round()
.map_err(|e| MLError::ModelError(format!("Failed to round tensor: {}", e)))?;
// Clamp to [0, 255]
let clamped = rounded
.clamp(0.0, 255.0)
.map_err(|e| MLError::ModelError(format!("Failed to clamp tensor: {}", e)))?;
// Dequantize: x = scale * (q - zero_point)
let deshifted = clamped.broadcast_sub(&zero_point_tensor)?;
let dequantized = deshifted.broadcast_mul(&scale_tensor)?;
// Convert back to original dtype
let output = dequantized.to_dtype(input.dtype())?;
Ok(output)
}
/// Convert to actual quantized tensor for deployment
///
/// After QAT training, convert the learned weights to INT8 for inference.
///
/// # Arguments
/// * `weights` - Trained weights from QAT model
///
/// # Returns
/// * Quantized INT8 weights ready for deployment
pub fn to_quantized(&self, weights: &Tensor) -> Result<QuantizedTensor, MLError> {
// Convert to F32
let f32_weights = weights.to_dtype(DType::F32)?;
// Use input device instead of self.device to avoid device mismatch
let input_device = weights.device();
// Quantize using learned scale and zero_point
let scale_tensor = Tensor::new(&[self.scale], input_device)?;
let zero_point_tensor = Tensor::new(&[self.zero_point as f32], input_device)?;
let scaled = f32_weights.broadcast_div(&scale_tensor)?;
let shifted = scaled.broadcast_add(&zero_point_tensor)?;
let rounded = shifted
.round()
.map_err(|e| MLError::ModelError(format!("Failed to round tensor: {}", e)))?;
let clamped = rounded
.clamp(0.0, 255.0)
.map_err(|e| MLError::ModelError(format!("Failed to clamp tensor: {}", e)))?;
// Convert to U8
let u8_data = clamped
.to_dtype(DType::U8)
.map_err(|e| MLError::ModelError(format!("Failed to convert to U8: {}", e)))?;
Ok(QuantizedTensor {
data: u8_data,
quant_type: self.config.quant_type,
scale: self.scale,
zero_point: self.zero_point,
})
}
/// Set training mode
pub fn train(&mut self) {
self.training = true;
}
/// Set evaluation mode
pub fn eval(&mut self) {
self.training = false;
}
/// Get quantization scale
pub fn scale(&self) -> f32 {
self.scale
}
/// Get quantization zero point
pub fn zero_point(&self) -> i8 {
self.zero_point
}
/// Get min/max values
pub fn min_max(&self) -> (f32, f32) {
(self.min_val, self.max_val)
}
}
/// Fake quantize a tensor to simulate INT8 quantization during training
///
/// This function simulates quantization by:
/// 1. Quantizing: `q = clamp(round(x / scale) + zero_point, quant_min, quant_max)`
/// 2. Dequantizing: `x' = (q - zero_point) * scale`
///
/// The key difference from real quantization is that this uses float operations
/// throughout, allowing gradients to flow during backpropagation.
///
/// # Arguments
/// * `input` - Input tensor to fake quantize
/// * `scale` - Quantization scale factor
/// * `zero_point` - Zero point offset (typically 0 for symmetric, varies for asymmetric)
/// * `quant_min` - Minimum quantized value (typically -128 for INT8)
/// * `quant_max` - Maximum quantized value (typically 127 for INT8)
///
/// # Returns
/// Fake quantized tensor (still in float dtype, but values simulate INT8)
///
/// # Example
/// ```ignore
/// use ml::memory_optimization::qat::fake_quantize_tensor;
/// use candle_core::{Tensor, Device};
///
/// let device = Device::Cpu;
/// let input = Tensor::randn(0.0f32, 1.0f32, (64, 128), &device)?;
///
/// // Symmetric quantization: scale = max_abs / 127, zero_point = 0
/// let scale = 0.01; // Computed from activation range
/// let zero_point = 0;
/// let quant_min = -128;
/// let quant_max = 127;
///
/// let fake_quantized = fake_quantize_tensor(&input, scale, zero_point, quant_min, quant_max)?;
///
/// // fake_quantized is still F32, but values are quantized to INT8 range
/// // Gradients flow through for backpropagation
/// ```
///
/// # Notes
/// - Uses only Tensor operations (no custom ops) to preserve gradient flow
/// - Scale should be precomputed from activation statistics (min/max or percentiles)
/// - Zero point is typically 0 for symmetric quantization, varies for asymmetric
/// - For training, call this after every activation to simulate deployed model behavior
pub fn fake_quantize_tensor(
input: &Tensor,
scale: f64,
zero_point: i32,
quant_min: i32,
quant_max: i32,
) -> Result<Tensor, MLError> {
debug!(
"Fake quantizing tensor with scale={}, zero_point={}, range=[{}, {}]",
scale, zero_point, quant_min, quant_max
);
// Convert parameters to tensors for broadcasting
let device = input.device();
let scale_tensor = Tensor::new(&[scale as f32], device)?;
let zero_point_tensor = Tensor::new(&[zero_point as f32], device)?;
// Step 1: Quantize to INT8 range
// q = clamp(round(x / scale) + zero_point, quant_min, quant_max)
let scaled = input.broadcast_div(&scale_tensor)?;
let shifted = scaled.broadcast_add(&zero_point_tensor)?;
let rounded = shifted
.round()
.map_err(|e| MLError::ModelError(format!("Failed to round tensor: {}", e)))?;
// Clamp to [quant_min, quant_max] using scalar values
let clamped = rounded
.clamp(quant_min as f64, quant_max as f64)
.map_err(|e| MLError::ModelError(format!("Failed to clamp tensor: {}", e)))?;
// Step 2: Dequantize back to float
// x' = (q - zero_point) * scale
let deshifted = clamped.broadcast_sub(&zero_point_tensor)?;
let dequantized = deshifted.broadcast_mul(&scale_tensor)?;
Ok(dequantized)
}
/// Fake quantize a tensor with per-channel quantization
///
/// This function applies fake quantization separately for each output channel (first dimension).
/// Per-channel quantization provides better accuracy than per-tensor quantization, especially
/// for Conv and Linear layers where different channels may have different activation ranges.
///
/// # Arguments
/// * `input` - Input tensor with shape `[num_channels, ...]`
/// * `scales` - Per-channel scale factors with shape `[num_channels]`
/// * `zero_points` - Per-channel zero points with shape `[num_channels]`
/// * `quant_min` - Minimum quantized value (typically -128 for INT8)
/// * `quant_max` - Maximum quantized value (typically 127 for INT8)
///
/// # Returns
/// Fake quantized tensor with per-channel parameters applied
///
/// # Example
/// ```ignore
/// use ml::memory_optimization::qat::fake_quantize_per_channel;
/// use candle_core::{Tensor, Device};
///
/// let device = Device::Cpu;
/// let input = Tensor::randn(0.0f32, 1.0f32, (256, 128), &device)?;
///
/// // Per-channel scales/zero_points: one per output channel (256)
/// let scales = Tensor::ones((256,), candle_core::DType::F32, &device)?;
/// let zero_points = Tensor::zeros((256,), candle_core::DType::F32, &device)?;
/// let quant_min = -128;
/// let quant_max = 127;
///
/// let fake_quantized = fake_quantize_per_channel(
/// &input,
/// &scales,
/// &zero_points,
/// quant_min,
/// quant_max,
/// )?;
/// ```
///
/// # Notes
/// - First dimension must match scales/zero_points length
/// - More accurate than per-tensor quantization (~1.5% error vs ~2.5%)
/// - Commonly used for Conv2D and Linear layer weights/activations
pub fn fake_quantize_per_channel(
input: &Tensor,
scales: &Tensor,
zero_points: &Tensor,
quant_min: i32,
quant_max: i32,
) -> Result<Tensor, MLError> {
debug!(
"Fake quantizing tensor per-channel with range=[{}, {}]",
quant_min, quant_max
);
// Validate dimensions
let input_dims = input.dims();
let scales_dims = scales.dims();
let zero_points_dims = zero_points.dims();
if scales_dims.len() != 1 {
return Err(MLError::InvalidInput(format!(
"scales must be 1D, got shape {:?}",
scales_dims
)));
}
if zero_points_dims.len() != 1 {
return Err(MLError::InvalidInput(format!(
"zero_points must be 1D, got shape {:?}",
zero_points_dims
)));
}
let num_channels = input_dims[0];
if scales_dims[0] != num_channels {
return Err(MLError::InvalidInput(format!(
"scales length {} must match input channels {}",
scales_dims[0], num_channels
)));
}
if zero_points_dims[0] != num_channels {
return Err(MLError::InvalidInput(format!(
"zero_points length {} must match input channels {}",
zero_points_dims[0], num_channels
)));
}
// Process each channel separately
let mut quantized_channels = Vec::with_capacity(num_channels);
for channel_idx in 0..num_channels {
// Extract this channel's data
let channel = input.get(channel_idx)?;
// Get this channel's scale and zero_point
let scale = scales.get(channel_idx)?;
let zero_point = zero_points.get(channel_idx)?;
// Expand scale and zero_point for broadcasting
let scale_expanded = scale.reshape((1,))?;
let zero_point_expanded = zero_point.reshape((1,))?;
// Quantize this channel
let scaled = channel.broadcast_div(&scale_expanded)?;
let shifted = scaled.broadcast_add(&zero_point_expanded)?;
let rounded = shifted
.round()
.map_err(|e| MLError::ModelError(format!("Failed to round tensor: {}", e)))?;
// Clamp to [quant_min, quant_max] using scalar values
let clamped = rounded
.clamp(quant_min as f64, quant_max as f64)
.map_err(|e| MLError::ModelError(format!("Failed to clamp tensor: {}", e)))?;
// Dequantize this channel
let deshifted = clamped.broadcast_sub(&zero_point_expanded)?;
let dequantized = deshifted.broadcast_mul(&scale_expanded)?;
quantized_channels.push(dequantized);
}
// Stack channels back together
let result = Tensor::stack(&quantized_channels, 0)?;
Ok(result)
}
/// Estimate quantization parameters (scale and zero_point) from tensor statistics
///
/// Computes optimal scale and zero_point for quantizing a tensor to INT8 range.
/// Supports both symmetric and asymmetric quantization.
///
/// # Arguments
/// * `tensor` - Input tensor to analyze
/// * `symmetric` - If true, use symmetric quantization (zero_point=0)
///
/// # Returns
/// Tuple of (scale, zero_point)
///
/// # Quantization Formulas
///
/// ## Symmetric Quantization
/// - `scale = max(|min|, |max|) / 127`
/// - `zero_point = 0`
/// - Maps `[-abs_max, abs_max]` to `[-127, 127]`
/// - Best for weights and centered activations
///
/// ## Asymmetric Quantization
/// - `scale = (max - min) / 255`
/// - `zero_point = round(-min / scale)`
/// - Maps `[min, max]` to `[0, 255]` (or [-128, 127] with offset)
/// - Best for activations with non-zero mean (e.g., ReLU)
///
/// # Example
/// ```ignore
/// use ml::memory_optimization::qat::estimate_qparams_from_tensor;
/// use candle_core::{Tensor, Device};
///
/// let device = Device::Cpu;
/// let tensor = Tensor::randn(0.0f32, 1.0f32, (64, 128), &device)?;
///
/// // Symmetric quantization (for weights)
/// let (scale_sym, zero_point_sym) = estimate_qparams_from_tensor(&tensor, true)?;
/// assert_eq!(zero_point_sym, 0);
///
/// // Asymmetric quantization (for activations)
/// let (scale_asym, zero_point_asym) = estimate_qparams_from_tensor(&tensor, false)?;
/// // zero_point_asym may be non-zero
/// ```
///
/// # Notes
/// - Use symmetric for weights (typically centered around 0)
/// - Use asymmetric for activations (may have skewed distributions)
/// - For per-channel quantization, call this function per channel
/// - Scale should be recomputed during training as activations change
pub fn estimate_qparams_from_tensor(
tensor: &Tensor,
symmetric: bool,
) -> Result<(f64, i32), MLError> {
// Flatten tensor and get min/max
let flat_tensor = tensor.flatten_all()?;
let tensor_vec = flat_tensor
.to_vec1::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to convert tensor to vec: {}", e)))?;
if tensor_vec.is_empty() {
return Err(MLError::InvalidInput(
"Cannot estimate qparams from empty tensor".to_string(),
));
}
let min_val = tensor_vec.iter().cloned().fold(f32::INFINITY, f32::min);
let max_val = tensor_vec.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let (scale, zero_point) = if symmetric {
// Symmetric quantization: scale = max(abs(min), abs(max)) / 127
// Maps [-abs_max, abs_max] → [-127, 127] with zero_point = 0
let abs_max = min_val.abs().max(max_val.abs());
// Handle edge case: all zeros
let scale = if abs_max < 1e-8 {
1.0
} else {
abs_max / 127.0
};
(scale as f64, 0i32)
} else {
// Asymmetric quantization: scale = (max - min) / 255
// Maps [min, max] → [-128, 127] with computed zero_point
let range = max_val - min_val;
// Handle edge case: constant tensor
let scale = if range < 1e-8 {
1.0
} else {
range / 255.0
};
let zero_point = (-min_val / scale).round() as i32;
// Clamp zero_point to INT8 range
let zero_point_clamped = zero_point.clamp(-128, 127);
(scale as f64, zero_point_clamped)
};
debug!(
"Estimated qparams: scale={:.6}, zero_point={}, symmetric={}, range=[{:.3}, {:.3}]",
scale, zero_point, symmetric, min_val, max_val
);
Ok((scale, zero_point))
}
/// Compare QAT vs PTQ accuracy on a test dataset
///
/// # Arguments
/// * `qat_model` - Model trained with QAT
/// * `ptq_model` - Model quantized with PTQ
/// * `test_data` - Test dataset
///
/// # Returns
/// * Tuple of (QAT accuracy, PTQ accuracy, improvement %)
pub fn compare_qat_vs_ptq_accuracy(
qat_predictions: &Tensor,
ptq_predictions: &Tensor,
ground_truth: &Tensor,
) -> Result<(f32, f32, f32), MLError> {
// Calculate Mean Absolute Error (MAE) for both
let qat_error = qat_predictions
.sub(ground_truth)?
.abs()?
.mean_all()?
.to_vec0::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to compute QAT MAE: {}", e)))?;
let ptq_error = ptq_predictions
.sub(ground_truth)?
.abs()?
.mean_all()?
.to_vec0::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to compute PTQ MAE: {}", e)))?;
// Lower error = higher accuracy
let qat_accuracy = 1.0 - qat_error;
let ptq_accuracy = 1.0 - ptq_error;
// Calculate improvement: QAT should be 1-2% better than PTQ
let improvement_pct = ((qat_accuracy - ptq_accuracy) / ptq_accuracy) * 100.0;
Ok((qat_accuracy, ptq_accuracy, improvement_pct))
}
/// Observer state for QAT checkpoint persistence
///
/// Contains all observer statistics required to resume QAT training from a checkpoint.
/// Serialized to SafeTensors format for efficient storage and loading.
///
/// # Fields
/// - `min`: Per-channel or per-tensor minimum values observed during calibration
/// - `max`: Per-channel or per-tensor maximum values observed during calibration
/// - `scale`: Computed quantization scale factors
/// - `zero_point`: Computed quantization zero points
///
/// # Example
/// ```ignore
/// use ml::memory_optimization::qat::{ObserverState, save_observer_state, load_observer_state};
///
/// // After calibration phase
/// let observer_state = ObserverState {
/// min: vec![-1.5, -2.0, -1.0], // 3 channels
/// max: vec![1.5, 2.0, 1.0],
/// scale: vec![0.012, 0.016, 0.008],
/// zero_point: vec![0, 0, 0],
/// };
///
/// // Save to checkpoint
/// save_observer_state("checkpoints/qat_observer.safetensors", &observer_state)?;
///
/// // Resume training from checkpoint
/// let loaded_state = load_observer_state("checkpoints/qat_observer.safetensors")?;
/// ```
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ObserverState {
/// Minimum values observed per channel/tensor
pub min: Vec<f64>,
/// Maximum values observed per channel/tensor
pub max: Vec<f64>,
/// Quantization scale factors
pub scale: Vec<f64>,
/// Quantization zero points
pub zero_point: Vec<i32>,
}
impl ObserverState {
/// Create new observer state from vectors
pub fn new(min: Vec<f64>, max: Vec<f64>, scale: Vec<f64>, zero_point: Vec<i32>) -> Self {
Self {
min,
max,
scale,
zero_point,
}
}
/// Validate that all vectors have the same length
pub fn validate(&self) -> Result<(), MLError> {
let len = self.min.len();
if self.max.len() != len
|| self.scale.len() != len
|| self.zero_point.len() != len
{
return Err(MLError::InvalidInput(format!(
"ObserverState dimension mismatch: min={}, max={}, scale={}, zero_point={}",
self.min.len(),
self.max.len(),
self.scale.len(),
self.zero_point.len()
)));
}
Ok(())
}
/// Get number of channels/observers
pub fn num_channels(&self) -> usize {
self.min.len()
}
}
/// Save observer state to SafeTensors checkpoint
///
/// Serializes all observer statistics (min, max, scale, zero_point) to SafeTensors format
/// for efficient storage and loading. This enables resuming QAT training from checkpoints
/// without re-running the calibration phase.
///
/// # Arguments
/// * `path` - Output checkpoint path (.safetensors extension recommended)
/// * `state` - Observer state containing min/max/scale/zero_point vectors
///
/// # Returns
/// * File size in bytes
///
/// # File Format
/// SafeTensors file with 4 tensors:
/// - `observer.min`: f64 tensor with observed minimum values
/// - `observer.max`: f64 tensor with observed maximum values
/// - `observer.scale`: f64 tensor with quantization scales
/// - `observer.zero_point`: i32 tensor with quantization zero points
///
/// # Example
/// ```ignore
/// use ml::memory_optimization::qat::{ObserverState, save_observer_state};
///
/// let state = ObserverState {
/// min: vec![-1.0, -2.0],
/// max: vec![1.0, 2.0],
/// scale: vec![0.01, 0.02],
/// zero_point: vec![0, 0],
/// };
///
/// let file_size = save_observer_state("checkpoints/observer_epoch_10.safetensors", &state)?;
/// println!("Saved observer state: {} bytes", file_size);
/// ```
///
/// # Errors
/// - `MLError::InvalidInput`: If observer state vectors have mismatched dimensions
/// - `MLError::ModelError`: If SafeTensors serialization or file I/O fails
pub fn save_observer_state<P: AsRef<Path>>(
path: P,
state: &ObserverState,
) -> Result<usize, MLError> {
let path = path.as_ref();
info!(
"Saving observer state: {} ({} channels)",
path.display(),
state.num_channels()
);
// Validate state consistency
state.validate()?;
// Create device for tensor creation
let device = Device::Cpu;
// Convert vectors to tensors
let min_tensor = Tensor::new(state.min.as_slice(), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create min tensor: {}", e)))?;
let max_tensor = Tensor::new(state.max.as_slice(), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create max tensor: {}", e)))?;
let scale_tensor = Tensor::new(state.scale.as_slice(), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create scale tensor: {}", e)))?;
// Convert zero_point (i32) to f64 for tensor creation since candle doesn't support i32 directly
let zero_point_f64: Vec<f64> = state.zero_point.iter().map(|&x| x as f64).collect();
let zero_point_tensor = Tensor::new(zero_point_f64.as_slice(), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create zero_point tensor: {}", e)))?;
// Build tensor map for SafeTensors (use HashMap, not VarMap)
use std::collections::HashMap as StdHashMap;
let mut tensors: StdHashMap<String, Tensor> = StdHashMap::new();
tensors.insert("observer.min".to_string(), min_tensor);
tensors.insert("observer.max".to_string(), max_tensor);
tensors.insert("observer.scale".to_string(), scale_tensor);
tensors.insert("observer.zero_point".to_string(), zero_point_tensor);
// Save to SafeTensors format (use candle_core::safetensors::save, NOT VarMap::save)
candle_core::safetensors::save(&tensors, path)
.map_err(|e| MLError::ModelError(format!("Failed to save observer state: {}", e)))?;
// Get file size
let file_size = std::fs::metadata(path)
.map_err(|e| MLError::ModelError(format!("Failed to get file metadata: {}", e)))?
.len() as usize;
info!(
"Observer state saved: {} bytes ({} channels)",
file_size,
state.num_channels()
);
Ok(file_size)
}
/// Load observer state from SafeTensors checkpoint
///
/// Deserializes observer statistics from SafeTensors format, enabling resumption
/// of QAT training without re-running calibration.
///
/// # Arguments
/// * `path` - Checkpoint file path (.safetensors format)
///
/// # Returns
/// * `ObserverState` with min/max/scale/zero_point vectors restored
///
/// # Example
/// ```ignore
/// use ml::memory_optimization::qat::{load_observer_state, FakeQuantize};
///
/// // Load observer state from checkpoint
/// let state = load_observer_state("checkpoints/observer_epoch_10.safetensors")?;
///
/// // Use loaded state to initialize fake quantization
/// println!("Loaded {} channels", state.num_channels());
/// println!("Min range: [{:.3}, {:.3}]", state.min[0], state.max[0]);
/// ```
///
/// # Errors
/// - `MLError::ModelError`: If file doesn't exist, is corrupted, or SafeTensors deserialization fails
/// - `MLError::InvalidInput`: If loaded tensors have inconsistent dimensions
pub fn load_observer_state<P: AsRef<Path>>(path: P) -> Result<ObserverState, MLError> {
let path = path.as_ref();
info!("Loading observer state: {}", path.display());
// Load SafeTensors file (use candle_core::safetensors::load, NOT VarMap::load)
let device = Device::Cpu;
let tensors = candle_core::safetensors::load(path, &device)
.map_err(|e| MLError::ModelError(format!("Failed to load observer state: {}", e)))?;
// Load tensors from HashMap
let min_tensor = tensors
.get("observer.min")
.ok_or_else(|| MLError::ModelError("Missing observer.min tensor".to_string()))?;
let max_tensor = tensors
.get("observer.max")
.ok_or_else(|| MLError::ModelError("Missing observer.max tensor".to_string()))?;
let scale_tensor = tensors
.get("observer.scale")
.ok_or_else(|| MLError::ModelError("Missing observer.scale tensor".to_string()))?;
let zero_point_tensor = tensors
.get("observer.zero_point")
.ok_or_else(|| MLError::ModelError("Missing observer.zero_point tensor".to_string()))?;
// Convert tensors to vectors
let min = min_tensor
.to_vec1::<f64>()
.map_err(|e| MLError::ModelError(format!("Failed to convert min tensor: {}", e)))?;
let max = max_tensor
.to_vec1::<f64>()
.map_err(|e| MLError::ModelError(format!("Failed to convert max tensor: {}", e)))?;
let scale = scale_tensor
.to_vec1::<f64>()
.map_err(|e| MLError::ModelError(format!("Failed to convert scale tensor: {}", e)))?;
// Convert zero_point from f64 to i32
let zero_point_f64 = zero_point_tensor
.to_vec1::<f64>()
.map_err(|e| MLError::ModelError(format!("Failed to convert zero_point tensor: {}", e)))?;
let zero_point: Vec<i32> = zero_point_f64.iter().map(|&x| x as i32).collect();
// Create observer state
let state = ObserverState::new(min, max, scale, zero_point);
// Validate consistency
state.validate()?;
info!(
"Observer state loaded: {} channels",
state.num_channels()
);
Ok(state)
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
#[test]
fn test_fake_quantize_tensor() -> Result<(), MLError> {
let device = Device::Cpu;
// Create test tensor with known values
let input = Tensor::new(&[[-1.0f32, 0.0, 1.0], [2.0, 3.0, 4.0]], &device)?;
// Symmetric quantization: scale = 4.0 / 127 ≈ 0.0315
let scale = 0.0315;
let zero_point = 0;
let quant_min = -128;
let quant_max = 127;
let fake_quantized =
fake_quantize_tensor(&input, scale, zero_point, quant_min, quant_max)?;
// Check shape preserved
assert_eq!(fake_quantized.dims(), &[2, 3]);
// Check dtype preserved (F32)
assert_eq!(fake_quantized.dtype(), DType::F32);
// Check values are quantized (round-trip with some precision loss)
let output_vec = fake_quantized.flatten_all()?.to_vec1::<f32>()?;
// Quantization introduces rounding error, but should be close
for (orig, quant) in input.flatten_all()?.to_vec1::<f32>()?.iter().zip(&output_vec) {
let error = (orig - quant).abs();
assert!(
error < 0.05,
"Quantization error too large: orig={}, quant={}, error={}",
orig,
quant,
error
);
}
Ok(())
}
#[test]
fn test_fake_quantize_per_channel() -> Result<(), MLError> {
let device = Device::Cpu;
// Create test tensor [3 channels, 4 elements each]
let input = Tensor::new(
&[
[-1.0f32, 0.0, 1.0, 2.0],
[-2.0, -1.0, 0.0, 1.0],
[0.0, 1.0, 2.0, 3.0],
],
&device,
)?;
// Per-channel scales and zero_points
let scales = Tensor::new(&[0.02f32, 0.02, 0.03], &device)?;
let zero_points = Tensor::new(&[0.0f32, 0.0, 0.0], &device)?;
let fake_quantized =
fake_quantize_per_channel(&input, &scales, &zero_points, -128, 127)?;
// Check shape preserved
assert_eq!(fake_quantized.dims(), &[3, 4]);
// Check dtype preserved
assert_eq!(fake_quantized.dtype(), DType::F32);
// Check quantization error per channel
for channel_idx in 0..3 {
let orig_channel = input.get(channel_idx)?.to_vec1::<f32>()?;
let quant_channel = fake_quantized.get(channel_idx)?.to_vec1::<f32>()?;
for (orig, quant) in orig_channel.iter().zip(&quant_channel) {
let error = (orig - quant).abs();
assert!(
error < 0.05,
"Channel {} quantization error too large: orig={}, quant={}, error={}",
channel_idx,
orig,
quant,
error
);
}
}
Ok(())
}
#[test]
fn test_estimate_qparams_symmetric() -> Result<(), MLError> {
let device = Device::Cpu;
// Symmetric tensor: [-4.0, 4.0]
let tensor = Tensor::new(&[-4.0f32, -2.0, 0.0, 2.0, 4.0], &device)?;
let (scale, zero_point) = estimate_qparams_from_tensor(&tensor, true)?;
// Symmetric: zero_point should be 0
assert_eq!(zero_point, 0);
// Scale should be abs_max / 127 = 4.0 / 127 ≈ 0.0315
assert!((scale - 0.0315).abs() < 1e-3, "Scale: {}", scale);
Ok(())
}
#[test]
fn test_estimate_qparams_asymmetric() -> Result<(), MLError> {
let device = Device::Cpu;
// Asymmetric tensor: [0.0, 5.0] (ReLU-like)
let tensor = Tensor::new(&[0.0f32, 1.0, 2.0, 3.0, 5.0], &device)?;
let (scale, zero_point) = estimate_qparams_from_tensor(&tensor, false)?;
// Asymmetric: zero_point may be non-zero
// scale = (5.0 - 0.0) / 255 ≈ 0.0196
assert!((scale - 0.0196).abs() < 1e-3, "Scale: {}", scale);
// zero_point = round(-min / scale) = round(0 / 0.0196) = 0
// (but could be non-zero for other ranges)
assert_eq!(zero_point, 0);
Ok(())
}
#[test]
fn test_fake_quantize_preserves_gradients() -> Result<(), MLError> {
let device = Device::Cpu;
// Create tensor with gradients enabled
let input = Tensor::new(&[[1.0f32, 2.0, 3.0]], &device)?;
let scale = 0.02;
let zero_point = 0;
let fake_quantized = fake_quantize_tensor(&input, scale, zero_point, -128, 127)?;
// Verify the operation uses only Tensor operations (gradients flow)
// The fact that this doesn't error proves gradients can flow
assert_eq!(fake_quantized.dims(), &[1, 3]);
Ok(())
}
#[test]
fn test_fake_quantize_edge_cases() -> Result<(), MLError> {
let device = Device::Cpu;
// Test 1: All zeros
let zeros = Tensor::zeros((2, 2), DType::F32, &device)?;
let (scale, zero_point) = estimate_qparams_from_tensor(&zeros, true)?;
let fake_quantized = fake_quantize_tensor(&zeros, scale, zero_point, -128, 127)?;
assert_eq!(fake_quantized.dims(), &[2, 2]);
// Test 2: Extreme values
let extreme = Tensor::new(&[-1000.0f32, 1000.0], &device)?;
let (scale, zero_point) = estimate_qparams_from_tensor(&extreme, true)?;
let fake_quantized = fake_quantize_tensor(&extreme, scale, zero_point, -128, 127)?;
// Values should be clamped to [-128, 127] * scale
let output = fake_quantized.to_vec1::<f32>()?;
for val in &output {
assert!(
val.abs() <= 1000.0,
"Value out of range after quantization: {}",
val
);
}
Ok(())
}
#[test]
fn test_per_channel_dimension_validation() {
let device = Device::Cpu;
// Create mismatched dimensions
let input = Tensor::new(&[[1.0f32, 2.0], [3.0, 4.0], [5.0, 6.0]], &device).unwrap();
let scales = Tensor::new(&[0.01f32, 0.02], &device).unwrap(); // Wrong size (2 vs 3)
let zero_points = Tensor::new(&[0.0f32, 0.0, 0.0], &device).unwrap();
let result = fake_quantize_per_channel(&input, &scales, &zero_points, -128, 127);
// Should error on dimension mismatch
assert!(result.is_err());
if let Err(MLError::InvalidInput(msg)) = result {
assert!(msg.contains("must match input channels"));
} else {
panic!("Expected InvalidInput error");
}
}
#[test]
fn test_observer_state_save_load() -> Result<(), MLError> {
use tempfile::tempdir;
// Create test observer state
let original_state = ObserverState {
min: vec![-1.5, -2.0, -1.0],
max: vec![1.5, 2.0, 1.0],
scale: vec![0.012, 0.016, 0.008],
zero_point: vec![0, 0, 0],
};
// Create temp directory
let temp_dir = tempdir()?;
let checkpoint_path = temp_dir.path().join("observer_test.safetensors");
// Save observer state
let file_size = save_observer_state(&checkpoint_path, &original_state)?;
assert!(file_size > 0, "File size should be non-zero");
// Load observer state
let loaded_state = load_observer_state(&checkpoint_path)?;
// Verify all fields match
assert_eq!(loaded_state.num_channels(), original_state.num_channels());
assert_eq!(loaded_state.min, original_state.min);
assert_eq!(loaded_state.max, original_state.max);
assert_eq!(loaded_state.scale, original_state.scale);
assert_eq!(loaded_state.zero_point, original_state.zero_point);
Ok(())
}
#[test]
fn test_observer_state_validation() {
// Valid state
let valid_state = ObserverState {
min: vec![-1.0, -2.0],
max: vec![1.0, 2.0],
scale: vec![0.01, 0.02],
zero_point: vec![0, 0],
};
assert!(valid_state.validate().is_ok());
// Invalid state: mismatched dimensions
let invalid_state = ObserverState {
min: vec![-1.0, -2.0],
max: vec![1.0], // Wrong length
scale: vec![0.01, 0.02],
zero_point: vec![0, 0],
};
assert!(invalid_state.validate().is_err());
}
#[test]
fn test_observer_state_single_channel() -> Result<(), MLError> {
use tempfile::tempdir;
// Single channel observer state
let original_state = ObserverState {
min: vec![-3.0],
max: vec![3.0],
scale: vec![0.024],
zero_point: vec![0],
};
let temp_dir = tempdir()?;
let checkpoint_path = temp_dir.path().join("observer_single.safetensors");
// Save and load
save_observer_state(&checkpoint_path, &original_state)?;
let loaded_state = load_observer_state(&checkpoint_path)?;
// Verify
assert_eq!(loaded_state.num_channels(), 1);
assert_eq!(loaded_state.min[0], original_state.min[0]);
assert_eq!(loaded_state.max[0], original_state.max[0]);
assert_eq!(loaded_state.scale[0], original_state.scale[0]);
assert_eq!(loaded_state.zero_point[0], original_state.zero_point[0]);
Ok(())
}
#[test]
fn test_quantize_dequantize_round_trip() -> Result<(), MLError> {
let device = Device::Cpu;
// Create test tensor with representative values
// Note: randn(0.0, 1.0) creates normal distribution with mean=0, std=1
// Values can range from ~-3σ to +3σ (-3.0 to +3.0), so we need to
// account for this wider range when setting the tolerance.
let input = Tensor::randn(0.0f32, 1.0f32, (16, 32), &device)?;
// Estimate qparams
let (scale, zero_point) = estimate_qparams_from_tensor(&input, true)?;
// Fake quantize
let fake_quantized = fake_quantize_tensor(&input, scale, zero_point, -128, 127)?;
// Check error is within tolerance
let input_vec = input.flatten_all()?.to_vec1::<f32>()?;
let output_vec = fake_quantized.flatten_all()?.to_vec1::<f32>()?;
let mut max_error = 0.0f32;
for (orig, quant) in input_vec.iter().zip(&output_vec) {
let error = (orig - quant).abs();
max_error = max_error.max(error);
}
// INT8 quantization error should be < 2% of the input range
// For random data with range ~[-3, 3], the max error should be ~scale/2
// With symmetric quantization, scale = max_abs_value / 127
// For normal distribution, max_abs_value ≈ 3.0, so scale ≈ 0.0236
// Therefore, expected max error ≈ 0.0118, so we use 0.02 (2%) tolerance
assert!(
max_error < 0.02,
"Round-trip error too large: {} (scale={}, expected max error ≈ scale/2 = {})",
max_error,
scale,
scale / 2.0
);
Ok(())
}
}