//! 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, DeviceLocation, DType, 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>>, /// Running maximum (EMA) running_max: Arc>>, /// Number of observations num_observations: Arc>, /// Calibration complete flag calibrated: Arc>, } 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::() .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 #[derive(Debug)] 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 { 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 { Ok(Self { config, device, scale, zero_point, min_val: 0.0, max_val: 255.0 * scale, training: true, }) } /// Check if two devices are the same (handles CUDA device IDs correctly) /// /// # CRITICAL FIX /// The original code used `std::mem::discriminant()` which only compared enum variant, /// NOT the contained data (CUDA ordinal). This caused silent device mismatches when /// comparing CUDA:0 vs CUDA:1. /// /// We also CANNOT use `Device::same_device()` or `CudaDevice::id()` because each call /// to `Device::cuda_if_available(0)` creates a NEW CudaDevice with a unique internal ID, /// even for the same CUDA ordinal. /// /// The correct approach is to use `Device::location()` which returns the actual CUDA ordinal. /// /// # Examples /// ```ignore /// let cuda0_a = Device::cuda_if_available(0)?; /// let cuda0_b = Device::cuda_if_available(0)?; /// let cuda1 = Device::cuda_if_available(1)?; /// /// assert!(FakeQuantize::devices_match(&cuda0_a, &cuda0_b)); // Same ordinal /// assert!(!FakeQuantize::devices_match(&cuda0_a, &cuda1)); // Different ordinals /// ``` fn devices_match(dev1: &Device, dev2: &Device) -> bool { match (dev1.location(), dev2.location()) { (DeviceLocation::Cpu, DeviceLocation::Cpu) => true, (DeviceLocation::Cuda { gpu_id: id1 }, DeviceLocation::Cuda { gpu_id: id2 }) => id1 == id2, (DeviceLocation::Metal { gpu_id: id1 }, DeviceLocation::Metal { gpu_id: id2 }) => id1 == id2, _ => false, // Different device types (CPU vs CUDA, etc.) } } /// 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, on FakeQuantize's device) /// /// # Device Handling /// Always ensures tensor operations happen on FakeQuantize's device to prevent /// CPU/CUDA mismatch. Input is moved to FakeQuantize device, processed, and /// returned on FakeQuantize's device (not original device). pub fn forward(&self, input: &Tensor) -> Result { if !self.training { // Evaluation mode: no quantization simulation return Ok(input.clone()); } // DEVICE CONSISTENCY FIX: Use proper device comparison instead of string formatting // Move input to FakeQuantize's device if needed (prevents CPU/CUDA mismatch) let input_on_device = if Self::devices_match(input.device(), &self.device) { // Same device (CPU or same CUDA ordinal): no move needed input.clone() } else { debug!( "Moving input from {:?} to {:?} for fake quantization", input.device(), &self.device ); input.to_device(&self.device)? }; // Convert to F32 for quantization let f32_input = input_on_device.to_dtype(DType::F32)?; // Quantize: q = clamp(round((x / scale) + zero_point), 0, 255) // All tensors now guaranteed on same device (self.device) let scale_tensor = Tensor::new(&[self.scale], &self.device)?; let zero_point_tensor = Tensor::new(&[self.zero_point as f32], &self.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())?; // DEVICE CONSISTENCY FIX: Always return tensor on FakeQuantize's device // Caller is responsible for moving to desired device if needed 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 { // 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) } /// Get scale and zero_point (for gradient clipping tests) pub fn scale_zero_point(&self) -> (f32, i8) { (self.scale, self.zero_point) } } /// 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 { 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 { 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::() .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::() .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::() .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, /// Maximum values observed per channel/tensor pub max: Vec, /// Quantization scale factors pub scale: Vec, /// Quantization zero points pub zero_point: Vec, } impl ObserverState { /// Create new observer state from vectors pub fn new(min: Vec, max: Vec, scale: Vec, zero_point: Vec) -> 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>( path: P, state: &ObserverState, ) -> Result { 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 = 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 = 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>(path: P) -> Result { 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::() .map_err(|e| MLError::ModelError(format!("Failed to convert min tensor: {}", e)))?; let max = max_tensor .to_vec1::() .map_err(|e| MLError::ModelError(format!("Failed to convert max tensor: {}", e)))?; let scale = scale_tensor .to_vec1::() .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::() .map_err(|e| MLError::ModelError(format!("Failed to convert zero_point tensor: {}", e)))?; let zero_point: Vec = 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::()?; // Quantization introduces rounding error, but should be close for (orig, quant) in input.flatten_all()?.to_vec1::()?.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::()?; let quant_channel = fake_quantized.get(channel_idx)?.to_vec1::()?; 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::()?; 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::()?; let output_vec = fake_quantized.flatten_all()?.to_vec1::()?; 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(()) } #[test] fn test_observer_checkpoint_round_trip() -> Result<(), MLError> { use tempfile::tempdir; println!("\n=== Test: Observer Checkpoint Round-Trip ==="); // Phase 1: Create and calibrate observer println!("📊 Phase 1: Calibrating observer with 100 batches"); let device = Device::Cpu; let config = QATConfig::default(); let mut observer = QuantizationObserver::new(config.clone(), device.clone()); // Calibrate with representative data for i in 0..100 { let batch = Tensor::randn(0f32, 1.0, (16, 256), &device)?; observer.observe(&batch)?; if i % 20 == 0 { println!(" ✓ Calibrated {}/100 batches", i + 1); } } assert!(observer.is_calibrated(), "Observer should be calibrated"); let (original_min, original_max) = observer.get_min_max().unwrap(); println!(" ✓ Calibration complete: min={:.6}, max={:.6}", original_min, original_max); // Phase 2: Create FakeQuantize from observer println!("📊 Phase 2: Creating FakeQuantize layer"); let original_fake_quant = FakeQuantize::from_observer(&observer)?; let original_scale = original_fake_quant.scale(); let original_zero_point = original_fake_quant.zero_point(); println!(" ✓ FakeQuantize created: scale={:.6}, zero_point={}", original_scale, original_zero_point); // Phase 3: Save observer state to checkpoint println!("📊 Phase 3: Saving observer state to checkpoint"); let temp_dir = tempdir()?; let checkpoint_path = temp_dir.path().join("observer_checkpoint.safetensors"); // Extract observer state let observer_state = ObserverState { min: vec![original_min as f64], max: vec![original_max as f64], scale: vec![original_scale as f64], zero_point: vec![original_zero_point as i32], }; let file_size = save_observer_state(&checkpoint_path, &observer_state)?; println!(" ✓ Checkpoint saved: {} bytes", file_size); assert!(file_size > 0, "Checkpoint file should have non-zero size"); // Phase 4: Load observer state from checkpoint println!("📊 Phase 4: Loading observer state from checkpoint"); let loaded_state = load_observer_state(&checkpoint_path)?; println!(" ✓ Observer state loaded: {} channels", loaded_state.num_channels()); // Validate loaded state loaded_state.validate()?; assert_eq!(loaded_state.num_channels(), 1, "Should have 1 channel"); println!(" ✓ Loaded state validated"); // Phase 5: Verify loaded statistics match original println!("📊 Phase 5: Verifying statistics match"); let loaded_min = loaded_state.min[0] as f32; let loaded_max = loaded_state.max[0] as f32; let loaded_scale = loaded_state.scale[0] as f32; let loaded_zero_point = loaded_state.zero_point[0] as i8; // Check min/max match let min_diff = (original_min - loaded_min).abs(); let max_diff = (original_max - loaded_max).abs(); assert!( min_diff < 1e-5, "Min mismatch: original={}, loaded={}, diff={}", original_min, loaded_min, min_diff ); assert!( max_diff < 1e-5, "Max mismatch: original={}, loaded={}, diff={}", original_max, loaded_max, max_diff ); println!(" ✓ Min/max match: min_diff={:.9}, max_diff={:.9}", min_diff, max_diff); // Check scale/zero_point match let scale_diff = (original_scale - loaded_scale).abs(); let zero_point_diff = (original_zero_point as i32 - loaded_zero_point as i32).abs(); assert!( scale_diff < 1e-5, "Scale mismatch: original={}, loaded={}, diff={}", original_scale, loaded_scale, scale_diff ); assert_eq!( original_zero_point, loaded_zero_point, "Zero point mismatch: original={}, loaded={}", original_zero_point, loaded_zero_point ); println!(" ✓ Scale/zero_point match: scale_diff={:.9}, zero_point_diff={}", scale_diff, zero_point_diff); // Phase 6: Validate numerical consistency println!("📊 Phase 6: Validating numerical consistency"); // Create new FakeQuantize from loaded state let loaded_fake_quant = FakeQuantize::new( config.clone(), device.clone(), loaded_scale, loaded_zero_point, )?; // Test with identical input let test_input = Tensor::new(&[[1.5f32, 2.3, -1.2, 0.5, 3.1, -2.8]], &device)?; let original_output = original_fake_quant.forward(&test_input)?; let loaded_output = loaded_fake_quant.forward(&test_input)?; // Compare outputs let original_data = original_output.flatten_all()?.to_vec1::()?; let loaded_data = loaded_output.flatten_all()?.to_vec1::()?; assert_eq!(original_data.len(), loaded_data.len(), "Output lengths should match"); let mut max_output_diff = 0.0f32; for (i, (orig_val, loaded_val)) in original_data.iter().zip(loaded_data.iter()).enumerate() { let diff = (orig_val - loaded_val).abs(); max_output_diff = max_output_diff.max(diff); assert!( diff < 1e-5, "Output mismatch at index {}: original={}, loaded={}, diff={}", i, orig_val, loaded_val, diff ); } println!(" ✓ Numerical consistency verified: max_output_diff={:.9}", max_output_diff); // Phase 7: Validate statistics are within expected ranges println!("📊 Phase 7: Validating statistics ranges"); // For standard normal distribution (mean=0, std=1): // - Expected min ≈ -3.0 (3σ below mean) // - Expected max ≈ +3.0 (3σ above mean) // - Expected scale ≈ 3.0 / 127 ≈ 0.024 // - Expected zero_point = 127 (symmetric quantization) assert!( loaded_min >= -4.0 && loaded_min <= 0.0, "Loaded min out of expected range: {} (expected [-4.0, 0.0])", loaded_min ); assert!( loaded_max >= 0.0 && loaded_max <= 4.0, "Loaded max out of expected range: {} (expected [0.0, 4.0])", loaded_max ); assert!( loaded_scale > 0.0 && loaded_scale < 0.1, "Loaded scale out of expected range: {} (expected (0.0, 0.1))", loaded_scale ); assert_eq!( loaded_zero_point, 127, "Loaded zero_point should be 127 for symmetric quantization, got {}", loaded_zero_point ); println!(" ✓ Statistics within expected ranges"); println!(" - Min: {:.6} (expected [-4.0, 0.0])", loaded_min); println!(" - Max: {:.6} (expected [0.0, 4.0])", loaded_max); println!(" - Scale: {:.6} (expected (0.0, 0.1))", loaded_scale); println!(" - Zero point: {} (expected 127)", loaded_zero_point); println!("✅ Observer checkpoint round-trip test passed!"); Ok(()) } #[test] fn test_devices_match_cpu() -> Result<(), MLError> { // Test that CPU devices always match let cpu1 = Device::Cpu; let cpu2 = Device::Cpu; assert!( FakeQuantize::devices_match(&cpu1, &cpu2), "CPU devices should match" ); Ok(()) } #[test] #[cfg(feature = "cuda")] fn test_devices_match_cuda_same_ordinal() -> Result<(), MLError> { // Test that CUDA devices with same ordinal match // This is the critical test that discriminant() would fail let cuda0_a = Device::cuda_if_available(0) .map_err(|e| MLError::ModelError(format!("CUDA not available: {}", e)))?; let cuda0_b = Device::cuda_if_available(0) .map_err(|e| MLError::ModelError(format!("CUDA not available: {}", e)))?; assert!( FakeQuantize::devices_match(&cuda0_a, &cuda0_b), "CUDA:0 should match CUDA:0 (different CudaDevice instances)" ); Ok(()) } #[test] #[cfg(feature = "cuda")] fn test_devices_match_cuda_different_ordinal() -> Result<(), MLError> { // Test that CUDA devices with different ordinals DO NOT match // This is the bug that discriminant() caused: CUDA:0 matched CUDA:1 let cuda0 = Device::cuda_if_available(0) .map_err(|e| MLError::ModelError(format!("CUDA not available: {}", e)))?; // Try to create CUDA:1 - may fail if only 1 GPU available match Device::cuda_if_available(1) { Ok(cuda1) => { assert!( !FakeQuantize::devices_match(&cuda0, &cuda1), "CUDA:0 should NOT match CUDA:1" ); } Err(_) => { // Only 1 GPU available, skip this test eprintln!("⚠️ Skipping CUDA:0 vs CUDA:1 test (only 1 GPU available)"); } } Ok(()) } #[test] fn test_devices_match_cpu_vs_cuda() -> Result<(), MLError> { // Test that CPU and CUDA devices do NOT match let cpu = Device::Cpu; #[cfg(feature = "cuda")] { if let Ok(cuda) = Device::cuda_if_available(0) { assert!( !FakeQuantize::devices_match(&cpu, &cuda), "CPU should NOT match CUDA" ); } } Ok(()) } #[test] fn test_fake_quantize_device_migration_cpu_to_cpu() -> Result<(), MLError> { // Test that FakeQuantize correctly handles CPU → CPU (no migration) let config = QATConfig::default(); let device = Device::Cpu; let fake_quant = FakeQuantize::new(config, device.clone(), 0.01, 0)?; // Create tensor on CPU let input = Tensor::new(&[[1.0f32, 2.0, 3.0]], &device)?; // Forward pass should work without device migration let output = fake_quant.forward(&input)?; // Verify output device assert!( FakeQuantize::devices_match(output.device(), &device), "Output should be on CPU" ); Ok(()) } #[test] #[cfg(feature = "cuda")] fn test_fake_quantize_device_migration_cuda_to_cuda() -> Result<(), MLError> { // Test that FakeQuantize correctly handles CUDA:0 → CUDA:0 (no migration) let config = QATConfig::default(); let device = Device::cuda_if_available(0) .map_err(|e| MLError::ModelError(format!("CUDA not available: {}", e)))?; let fake_quant = FakeQuantize::new(config, device.clone(), 0.01, 0)?; // Create tensor on CUDA:0 let input = Tensor::new(&[[1.0f32, 2.0, 3.0]], &device)?; // Forward pass should work without device migration let output = fake_quant.forward(&input)?; // Verify output device matches input (CUDA:0) assert!( FakeQuantize::devices_match(output.device(), &device), "Output should be on CUDA:0" ); Ok(()) } #[test] #[cfg(feature = "cuda")] fn test_fake_quantize_device_migration_cpu_to_cuda() -> Result<(), MLError> { // Test that FakeQuantize correctly migrates CPU tensor to CUDA device let config = QATConfig::default(); let cuda_device = Device::cuda_if_available(0) .map_err(|e| MLError::ModelError(format!("CUDA not available: {}", e)))?; let fake_quant = FakeQuantize::new(config, cuda_device.clone(), 0.01, 0)?; // Create tensor on CPU (device mismatch) let cpu_device = Device::Cpu; let input = Tensor::new(&[[1.0f32, 2.0, 3.0]], &cpu_device)?; // Forward pass should migrate CPU → CUDA automatically let output = fake_quant.forward(&input)?; // Verify output device is CUDA (FakeQuantize's device) assert!( FakeQuantize::devices_match(output.device(), &cuda_device), "Output should be on CUDA after migration from CPU" ); Ok(()) } }