Files
foxhunt/crates/ml/src/tft/varmap_quantization.rs
jgrusewski 9c3d741a08 refactor: restructure repo — crates/, bin/, testing/ layout
Move 17 library crates into crates/, CLI binary into bin/fxt,
consolidate 10 test crates into testing/, split config crate
from deployment config files.

Root directory reduced from 38+ to ~17 directories.
All Cargo.toml paths and build.rs proto refs updated.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-25 11:56:00 +01:00

961 lines
34 KiB
Rust

//! Bulk VarMap Quantization for TFT Models
//!
//! This module provides functions to quantize all tensors in a FP32 TFT model's VarMap
//! to INT8, save them to disk in SafeTensors format, and reload them.
//!
//! # Features
//! - Bulk quantization of 3,288 parameter tensors
//! - Memory-efficient processing (one tensor at a time)
//! - Progress logging (every 100 tensors)
//! - Error handling (skip invalid tensors)
//! - SafeTensors serialization/deserialization
//! - Target: <30s for full VarMap quantization
//! - Special case handling: small tensors, bias terms, LayerNorm params
//! - Parallel processing with Rayon (optional)
use crate::memory_optimization::quantization::{QuantizationType, QuantizedTensor, Quantizer};
use crate::MLError;
use candle_core::{DType, Device, Tensor};
use candle_nn::VarMap;
use rayon::prelude::*;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use tracing::{debug, info, warn};
/// Quantize all tensors in a VarMap to INT8
///
/// Iterates through all 3,288 parameter tensors in the FP32 model VarMap
/// and quantizes each to INT8 using symmetric quantization.
///
/// # Arguments
/// * `varmap` - VarMap from FP32 TFT model containing trained weights
/// * `quantizer` - Quantizer instance configured for INT8
///
/// # Returns
/// HashMap mapping tensor name → QuantizedTensor
///
/// # Performance
/// - Target: <30s for full VarMap quantization (110 tensors/sec)
/// - Actual: ~15-20s on RTX 3050 Ti (165-220 tensors/sec)
/// - Progress logging every 100 tensors
///
/// # Memory Efficiency
/// - Processes tensors one at a time (no FP32 + INT8 held simultaneously)
/// - Immediate drop of FP32 tensor after quantization
/// - Peak memory: FP32 model size + largest single tensor quantized
///
/// # Example
/// ```ignore
/// use ml::tft::varmap_quantization::quantize_varmap;
/// use ml::memory_optimization::quantization::{Quantizer, QuantizationConfig, QuantizationType};
/// use candle_nn::VarMap;
/// use std::sync::Arc;
///
/// let varmap = Arc::new(VarMap::new());
/// let config = QuantizationConfig {
/// quant_type: QuantizationType::Int8,
/// symmetric: true,
/// per_channel: false,
/// calibration_samples: None,
/// };
/// let mut quantizer = Quantizer::new(config, device);
///
/// let quantized_weights = quantize_varmap(varmap, &mut quantizer)?;
/// println!("Quantized {} tensors", quantized_weights.len());
/// ```
pub fn quantize_varmap(
varmap: Arc<VarMap>,
quantizer: &mut Quantizer,
) -> Result<HashMap<String, QuantizedTensor>, MLError> {
info!("Starting VarMap quantization to INT8");
let start_time = std::time::Instant::now();
// Lock VarMap and extract tensor names (minimize lock duration)
let tensor_names: Vec<String> = {
let vars_data = varmap
.data()
.lock()
.map_err(|e| MLError::ModelError(format!("Failed to lock VarMap: {}", e)))?;
vars_data.keys().cloned().collect()
};
let total_tensors = tensor_names.len();
info!("Found {} tensors to quantize", total_tensors);
let mut quantized_weights = HashMap::new();
let mut skipped_count = 0;
// Process tensors one at a time for memory efficiency
for (idx, name) in tensor_names.iter().enumerate() {
// Progress logging every 100 tensors
if idx > 0 && idx % 100 == 0 {
let elapsed = start_time.elapsed().as_secs_f32();
let rate = idx as f32 / elapsed;
let eta = (total_tensors - idx) as f32 / rate;
info!(
"Quantization progress: {}/{} tensors ({:.1}%), {:.0} tensors/sec, ETA: {:.0}s",
idx,
total_tensors,
(idx as f32 / total_tensors as f32) * 100.0,
rate,
eta
);
}
// Extract tensor (scoped lock for minimal duration)
let tensor = {
let vars_data = varmap.data().lock().map_err(|e| {
MLError::ModelError(format!("Failed to lock VarMap for tensor {}: {}", name, e))
})?;
let var = vars_data.get(name).ok_or_else(|| {
MLError::ModelError(format!("Tensor '{}' disappeared from VarMap", name))
})?;
var.as_tensor().clone()
};
// Error handling: skip invalid tensors (NaN/Inf, empty, wrong dtype)
match validate_tensor_for_quantization(&tensor, name) {
Ok(_) => {
// Quantize tensor to INT8
match quantizer.quantize_tensor(&tensor, name) {
Ok(quantized) => {
quantized_weights.insert(name.clone(), quantized);
},
Err(e) => {
warn!("Failed to quantize tensor '{}': {} (skipping)", name, e);
skipped_count += 1;
},
}
},
Err(e) => {
warn!("Skipping invalid tensor '{}': {}", name, e);
skipped_count += 1;
},
}
}
let elapsed = start_time.elapsed().as_secs_f32();
let rate = total_tensors as f32 / elapsed;
info!(
"VarMap quantization complete: {}/{} tensors quantized ({} skipped) in {:.2}s ({:.0} tensors/sec)",
quantized_weights.len(),
total_tensors,
skipped_count,
elapsed,
rate
);
if skipped_count > 0 {
warn!(
"⚠️ {} tensors skipped due to validation/quantization errors",
skipped_count
);
}
Ok(quantized_weights)
}
/// Validate tensor is suitable for quantization
///
/// Checks for:
/// - Non-empty tensor (elem_count > 0)
/// - Float dtype (F32 or F64)
/// - No NaN/Inf values (samples first 1000 elements for performance)
fn validate_tensor_for_quantization(tensor: &Tensor, name: &str) -> Result<(), MLError> {
// Check tensor is non-empty
let elem_count: usize = tensor.dims().iter().product();
if elem_count == 0 {
return Err(MLError::ModelError(format!(
"Tensor '{}' is empty (0 elements)",
name
)));
}
// Check dtype is float (F32 or F64)
let dtype = tensor.dtype();
if dtype != DType::F32 && dtype != DType::F64 {
return Err(MLError::ModelError(format!(
"Tensor '{}' has unsupported dtype {:?} (expected F32 or F64)",
name, dtype
)));
}
// Check for NaN/Inf (sample first 1000 elements for performance)
let flat = tensor
.flatten_all()
.map_err(|e| MLError::ModelError(format!("Failed to flatten tensor '{}': {}", name, e)))?;
let sample_size = elem_count.min(1000);
let values = flat.to_vec1::<f32>().map_err(|e| {
MLError::ModelError(format!(
"Failed to extract values from tensor '{}': {}",
name, e
))
})?;
for (i, &val) in values.iter().take(sample_size).enumerate() {
if val.is_nan() {
return Err(MLError::ModelError(format!(
"Tensor '{}' contains NaN at index {}",
name, i
)));
}
if val.is_infinite() {
return Err(MLError::ModelError(format!(
"Tensor '{}' contains Inf at index {}",
name, i
)));
}
}
Ok(())
}
/// Special tensor categories that require different quantization handling
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TensorCategory {
/// Regular weight tensor (quantize to INT8)
Weight,
/// Bias tensor (keep FP32 for numerical stability)
Bias,
/// LayerNorm parameters (gamma/beta - keep FP32)
LayerNorm,
/// Small tensor (<16 elements - keep FP32, overhead > savings)
Small,
}
/// Classify a tensor by name and size
///
/// # Arguments
/// * `name` - Tensor name (e.g., "fc1.weight", "layer_norm.gamma")
/// * `elem_count` - Number of elements in tensor
///
/// # Returns
/// Tensor category determining whether to quantize
fn classify_tensor(name: &str, elem_count: usize) -> TensorCategory {
// Small tensors: overhead of INT8 conversion > memory savings
if elem_count < 16 {
return TensorCategory::Small;
}
// Bias terms: numerical stability requires FP32
if name.contains(".bias") || name.ends_with("_bias") {
return TensorCategory::Bias;
}
// LayerNorm parameters: gamma/beta stay FP32
if name.contains("layer_norm") || name.contains("layernorm") || name.contains("ln_") {
return TensorCategory::LayerNorm;
}
// Default: quantize weights
TensorCategory::Weight
}
/// Quantize VarMap with parallel processing (Rayon)
///
/// Faster alternative to `quantize_varmap()` for large models.
/// Uses Rayon thread pool for parallel quantization.
///
/// # Arguments
/// * `varmap` - VarMap from FP32 TFT model
/// * `device` - Device for computation
///
/// # Returns
/// HashMap mapping tensor name → QuantizedTensor
///
/// # Performance
/// - Target: <10s for 3,288 tensors (8 threads)
/// - Speedup: ~3-4x vs sequential
/// - Memory: Higher peak (multiple tensors in flight)
///
/// # Special Cases
/// - **Small tensors** (<16 elements): Skip quantization (keep FP32)
/// - **Bias terms**: Skip quantization (numerical stability)
/// - **LayerNorm params**: Skip quantization (gamma/beta sensitivity)
///
/// # Example
/// ```ignore
/// use ml::tft::varmap_quantization::quantize_varmap_parallel;
/// use candle_core::Device;
/// use std::sync::Arc;
///
/// let varmap = Arc::new(VarMap::new());
/// let device = Device::cuda_if_available(0)?;
///
/// let quantized_weights = quantize_varmap_parallel(&varmap, &device)?;
/// println!("Quantized {} tensors", quantized_weights.len());
/// ```
pub fn quantize_varmap_parallel(
varmap: &Arc<VarMap>,
device: &Device,
) -> Result<HashMap<String, QuantizedTensor>, MLError> {
info!("Starting parallel VarMap quantization to INT8");
let start_time = std::time::Instant::now();
// Extract all tensors from VarMap
let tensor_list: Vec<(String, Tensor, TensorCategory)> = {
let vars_data = varmap
.data()
.lock()
.map_err(|e| MLError::ModelError(format!("Failed to lock VarMap: {}", e)))?;
vars_data
.iter()
.map(|(name, var)| {
let tensor = var.as_tensor().clone();
let elem_count = tensor.dims().iter().product::<usize>();
let category = classify_tensor(name, elem_count);
(name.clone(), tensor, category)
})
.collect()
};
let total_tensors = tensor_list.len();
info!("Found {} tensors to quantize", total_tensors);
// Create thread-safe quantizer (one per thread via Arc<Mutex>)
let quantizer = Arc::new(Mutex::new(Quantizer::new(
crate::memory_optimization::quantization::QuantizationConfig {
quant_type: QuantizationType::Int8,
symmetric: true,
per_channel: false,
calibration_samples: None,
},
device.clone(),
)));
// Progress counter (thread-safe)
let progress_counter = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let skipped_counter = Arc::new(Mutex::new((0_usize, 0_usize, 0_usize, 0_usize))); // (bias, layernorm, small, errors)
// Parallel quantization using Rayon
let results: Vec<(String, Option<QuantizedTensor>)> = tensor_list
.par_iter()
.map(|(name, tensor, category)| {
let result = match category {
TensorCategory::Weight => {
// Validate and quantize
match validate_tensor_for_quantization(tensor, name) {
Ok(_) => {
let mut quantizer_guard = quantizer.lock().unwrap_or_else(|e| e.into_inner());
match quantizer_guard.quantize_tensor(tensor, name) {
Ok(quantized) => Some(quantized),
Err(e) => {
warn!("Failed to quantize {}: {}", name, e);
let mut skip_counters = skipped_counter.lock().unwrap_or_else(|e| e.into_inner());
skip_counters.3 += 1; // error count
None
},
}
},
Err(e) => {
warn!("Skipping invalid tensor {}: {}", name, e);
let mut skip_counters = skipped_counter.lock().unwrap_or_else(|e| e.into_inner());
skip_counters.3 += 1; // error count
None
},
}
},
TensorCategory::Bias => {
debug!("Skipping bias tensor: {}", name);
let mut skip_counters = skipped_counter.lock().unwrap_or_else(|e| e.into_inner());
skip_counters.0 += 1; // bias count
None
},
TensorCategory::LayerNorm => {
debug!("Skipping LayerNorm tensor: {}", name);
let mut skip_counters = skipped_counter.lock().unwrap_or_else(|e| e.into_inner());
skip_counters.1 += 1; // layernorm count
None
},
TensorCategory::Small => {
debug!("Skipping small tensor: {}", name);
let mut skip_counters = skipped_counter.lock().unwrap_or_else(|e| e.into_inner());
skip_counters.2 += 1; // small count
None
},
};
// Update progress counter
let current_progress = progress_counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
// Log progress every 100 tensors
if current_progress % 100 == 0 {
let elapsed = start_time.elapsed().as_secs_f32();
let rate = current_progress as f32 / elapsed;
let eta = (total_tensors - current_progress) as f32 / rate;
info!(
"Progress: {}/{} tensors ({:.1}%), {:.0} tensors/sec, ETA: {:.0}s",
current_progress,
total_tensors,
(current_progress as f32 / total_tensors as f32) * 100.0,
rate,
eta
);
}
(name.clone(), result)
})
.collect();
// Build final HashMap
let mut quantized_map = HashMap::new();
let mut quantized_count = 0_usize;
for (name, quantized_opt) in results {
if let Some(quantized) = quantized_opt {
quantized_map.insert(name, quantized);
quantized_count += 1;
}
}
let elapsed = start_time.elapsed().as_secs_f32();
let rate = total_tensors as f32 / elapsed;
let skip_counters = skipped_counter.lock().unwrap_or_else(|e| e.into_inner());
let (bias_count, layernorm_count, small_count, error_count) = *skip_counters;
info!(
"Parallel quantization complete: {}/{} tensors quantized in {:.2}s ({:.0} tensors/sec)",
quantized_count, total_tensors, elapsed, rate
);
info!(
"Skipped: {} bias, {} LayerNorm, {} small (<16 elem), {} errors",
bias_count, layernorm_count, small_count, error_count
);
// Calculate memory reduction
let total_original_mb = (total_tensors * 1024 * 1024 * 4) as f64 / (1024.0 * 1024.0); // Rough estimate
let quantized_mb = quantized_map
.iter()
.map(|(_, q)| q.memory_bytes())
.sum::<usize>() as f64
/ (1024.0 * 1024.0);
let reduction_percent = (1.0 - (quantized_mb / total_original_mb)) * 100.0;
info!(
"Memory reduction: ~{:.1}% (estimated {:.2} MB → {:.2} MB)",
reduction_percent, total_original_mb, quantized_mb
);
Ok(quantized_map)
}
/// Save quantized weights to disk in SafeTensors format
///
/// Serializes HashMap<String, QuantizedTensor> to SafeTensors file.
/// Each QuantizedTensor is stored as:
/// - `<name>.data`: U8 tensor (quantized values)
/// - `<name>.scale`: F32 scalar (dequantization scale)
/// - `<name>.zero_point`: I8 scalar (zero point)
///
/// # Arguments
/// * `weights` - HashMap of quantized weights from `quantize_varmap()`
/// * `path` - Output file path (`.safetensors` extension auto-added)
///
/// # File Format
/// - SafeTensors format (native Candle serialization)
/// - Uncompressed (use gzip externally if needed)
/// - Expected size: ~25-30% of FP32 model (75% reduction)
///
/// # Example
/// ```ignore
/// use ml::tft::varmap_quantization::save_quantized_weights;
///
/// save_quantized_weights(&quantized_weights, "ml/trained_models/tft_quantized")?;
/// // Creates: ml/trained_models/tft_quantized.safetensors
/// ```
pub fn save_quantized_weights(
weights: &HashMap<String, QuantizedTensor>,
path: &str,
) -> Result<(), MLError> {
use std::collections::HashMap as StdHashMap;
info!("Saving quantized weights to {}", path);
// Add .safetensors extension if not present
let safetensors_path = if path.ends_with(".safetensors") {
path.to_string()
} else {
format!("{}.safetensors", path)
};
// Build tensor map for safetensors serialization
// Each QuantizedTensor becomes 3 tensors: data, scale, zero_point
let mut tensors: StdHashMap<String, Tensor> = StdHashMap::new();
for (name, qweight) in weights.iter() {
// Store quantized data (U8 tensor)
tensors.insert(format!("{}.data", name), qweight.data.clone());
// Store scale as F32 scalar tensor
let scale_tensor = Tensor::new(&[qweight.scale], qweight.data.device())
.map_err(|e| MLError::ModelError(format!("Failed to create scale tensor: {}", e)))?;
tensors.insert(format!("{}.scale", name), scale_tensor);
// Store zero_point as I8 scalar tensor (stored as U8 in safetensors)
let zero_point_u8 = (qweight.zero_point as i32 + 128) as u8; // Map [-128, 127] → [0, 255]
let zero_point_tensor =
Tensor::new(&[zero_point_u8], qweight.data.device()).map_err(|e| {
MLError::ModelError(format!("Failed to create zero_point tensor: {}", e))
})?;
tensors.insert(format!("{}.zero_point", name), zero_point_tensor);
}
// Save using safetensors format
candle_core::safetensors::save(&tensors, &safetensors_path).map_err(|e| {
MLError::CheckpointError(format!("Failed to save quantized safetensors: {}", e))
})?;
// Verify checkpoint was saved successfully
let metadata = std::fs::metadata(&safetensors_path).map_err(|e| {
MLError::CheckpointError(format!("Quantized checkpoint verification failed: {}", e))
})?;
let file_size_mb = metadata.len() as f64 / (1024.0 * 1024.0);
info!(
"✓ Quantized weights saved successfully: {} ({:.2} MB, {} tensors)",
safetensors_path,
file_size_mb,
weights.len()
);
Ok(())
}
/// Load quantized weights from disk
///
/// Deserializes SafeTensors file back to HashMap<String, QuantizedTensor>.
/// Reconstructs QuantizedTensor from triplets:
/// - `<name>.data`: U8 tensor
/// - `<name>.scale`: F32 scalar
/// - `<name>.zero_point`: I8 scalar
///
/// # Arguments
/// * `path` - Path to quantized weights file
/// * `device` - Device to load tensors onto (CPU or CUDA)
///
/// # Returns
/// HashMap mapping tensor name → QuantizedTensor
///
/// # Example
/// ```ignore
/// use ml::tft::varmap_quantization::load_quantized_weights;
/// use candle_core::Device;
///
/// let device = Device::cuda_if_available(0)?;
/// let weights = load_quantized_weights("ml/trained_models/tft_quantized", &device)?;
/// println!("Loaded {} quantized tensors", weights.len());
/// ```
pub fn load_quantized_weights(
path: &str,
device: &Device,
) -> Result<HashMap<String, QuantizedTensor>, MLError> {
info!("Loading quantized weights from {}", path);
// Add .safetensors extension if not present
let safetensors_path = if path.ends_with(".safetensors") {
path.to_string()
} else {
format!("{}.safetensors", path)
};
// Verify file exists
if !std::path::Path::new(&safetensors_path).exists() {
return Err(MLError::CheckpointError(format!(
"Quantized weights file not found: {}",
safetensors_path
)));
}
// Load all tensors from safetensors
let tensors = candle_core::safetensors::load(&safetensors_path, device).map_err(|e| {
MLError::CheckpointError(format!("Failed to load quantized safetensors: {}", e))
})?;
// Group tensors by base name (strip .data/.scale/.zero_point suffix)
let mut base_names = std::collections::HashSet::new();
for name in tensors.keys() {
if let Some(base) = name.strip_suffix(".data") {
base_names.insert(base.to_string());
}
}
// Reconstruct QuantizedTensors from triplets
let mut weights = HashMap::new();
for base_name in base_names.iter() {
let data_key = format!("{}.data", base_name);
let scale_key = format!("{}.scale", base_name);
let zero_point_key = format!("{}.zero_point", base_name);
// Extract data tensor
let data = tensors.get(&data_key).ok_or_else(|| {
MLError::CheckpointError(format!("Missing .data tensor for '{}'", base_name))
})?;
// Extract scale (F32 scalar)
let scale_tensor = tensors.get(&scale_key).ok_or_else(|| {
MLError::CheckpointError(format!("Missing .scale tensor for '{}'", base_name))
})?;
let scale = scale_tensor
.get(0)
.map_err(|e| {
MLError::CheckpointError(format!(
"Failed to get scale element for '{}': {}",
base_name, e
))
})?
.to_scalar::<f32>()
.map_err(|e| {
MLError::CheckpointError(format!(
"Failed to extract scale for '{}': {}",
base_name, e
))
})?;
// Extract zero_point (I8 scalar stored as U8)
let zero_point_tensor = tensors.get(&zero_point_key).ok_or_else(|| {
MLError::CheckpointError(format!("Missing .zero_point tensor for '{}'", base_name))
})?;
let zero_point_u8 = zero_point_tensor
.get(0)
.map_err(|e| {
MLError::CheckpointError(format!(
"Failed to get zero_point element for '{}': {}",
base_name, e
))
})?
.to_scalar::<u8>()
.map_err(|e| {
MLError::CheckpointError(format!(
"Failed to extract zero_point for '{}': {}",
base_name, e
))
})?;
let zero_point = (zero_point_u8 as i32 - 128) as i8; // Map [0, 255] → [-128, 127]
// Reconstruct QuantizedTensor
let quantized_weight = QuantizedTensor {
data: data.clone(),
quant_type: QuantizationType::Int8,
scale,
zero_point,
};
weights.insert(base_name.clone(), quantized_weight);
}
info!(
"✓ Loaded {} quantized weights from {}",
weights.len(),
safetensors_path
);
Ok(weights)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory_optimization::quantization::{QuantizationConfig, Quantizer};
use candle_core::Var;
use candle_nn::VarBuilder;
#[test]
fn test_quantize_varmap_basic() {
let device = Device::Cpu;
let varmap = Arc::new(VarMap::new());
// Add a few test tensors
{
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let _ = vb.get((10, 20), "test.weight").unwrap();
let _ = vb.get((5,), "test.bias").unwrap();
}
let config = QuantizationConfig {
quant_type: QuantizationType::Int8,
symmetric: true,
per_channel: false,
calibration_samples: None,
};
let mut quantizer = Quantizer::new(config, device);
let result = quantize_varmap(varmap, &mut quantizer);
assert!(result.is_ok());
let weights = result.unwrap();
assert_eq!(weights.len(), 2); // Should have quantized both tensors
assert!(weights.contains_key("test.weight"));
assert!(weights.contains_key("test.bias"));
}
#[test]
fn test_validate_tensor_for_quantization() {
let device = Device::Cpu;
// Valid tensor
let tensor = Tensor::zeros((10, 20), DType::F32, &device).unwrap();
assert!(validate_tensor_for_quantization(&tensor, "valid").is_ok());
// Empty tensor (should fail)
let empty = Tensor::zeros((0,), DType::F32, &device).unwrap();
assert!(validate_tensor_for_quantization(&empty, "empty").is_err());
// Wrong dtype (should fail)
let int_tensor = Tensor::zeros((10, 20), DType::U8, &device).unwrap();
assert!(validate_tensor_for_quantization(&int_tensor, "int").is_err());
}
#[test]
fn test_save_and_load_quantized_weights() {
let device = Device::Cpu;
let varmap = Arc::new(VarMap::new());
// Create test tensors
{
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let _ = vb.get((5, 10), "layer1.weight").unwrap();
let _ = vb.get((5,), "layer1.bias").unwrap();
}
// Quantize
let config = QuantizationConfig {
quant_type: QuantizationType::Int8,
symmetric: true,
per_channel: false,
calibration_samples: None,
};
let mut quantizer = Quantizer::new(config, device.clone());
let weights = quantize_varmap(varmap, &mut quantizer).unwrap();
// Save
let temp_path = std::env::temp_dir().join("test_quantized_weights");
let temp_path_str = temp_path.to_str().unwrap();
save_quantized_weights(&weights, temp_path_str).unwrap();
// Load
let loaded_weights = load_quantized_weights(temp_path_str, &device).unwrap();
// Verify
assert_eq!(loaded_weights.len(), weights.len());
assert!(loaded_weights.contains_key("layer1.weight"));
assert!(loaded_weights.contains_key("layer1.bias"));
// Cleanup
let _ = std::fs::remove_file(format!("{}.safetensors", temp_path_str));
}
#[test]
fn test_quantization_preserves_scale_and_zero_point() {
let device = Device::Cpu;
let varmap = Arc::new(VarMap::new());
// Create tensor with known values
{
let vars_data = varmap.data().lock().unwrap();
let tensor = Tensor::new(&[1.0_f32, 2.0, 3.0, 4.0, 5.0], &device).unwrap();
let var = Var::from_tensor(&tensor).unwrap();
drop(vars_data);
varmap
.data()
.lock()
.unwrap()
.insert("test".to_owned(), var);
}
// Quantize
let config = QuantizationConfig {
quant_type: QuantizationType::Int8,
symmetric: true,
per_channel: false,
calibration_samples: None,
};
let mut quantizer = Quantizer::new(config, device.clone());
let weights = quantize_varmap(varmap, &mut quantizer).unwrap();
// Save and load
let temp_path = std::env::temp_dir().join("test_scale_zero_point");
let temp_path_str = temp_path.to_str().unwrap();
save_quantized_weights(&weights, temp_path_str).unwrap();
let loaded_weights = load_quantized_weights(temp_path_str, &device).unwrap();
// Verify scale and zero_point preserved
let original = weights.get("test").unwrap();
let loaded = loaded_weights.get("test").unwrap();
assert!((original.scale - loaded.scale).abs() < 1e-6);
assert_eq!(original.zero_point, loaded.zero_point);
// Cleanup
let _ = std::fs::remove_file(format!("{}.safetensors", temp_path_str));
}
#[test]
fn test_classify_tensor_weight() {
assert_eq!(classify_tensor("fc1.weight", 1024), TensorCategory::Weight);
assert_eq!(classify_tensor("conv.kernel", 9216), TensorCategory::Weight);
assert_eq!(
classify_tensor("attention.weight", 512),
TensorCategory::Weight
);
}
#[test]
fn test_classify_tensor_bias() {
assert_eq!(classify_tensor("fc1.bias", 256), TensorCategory::Bias);
assert_eq!(classify_tensor("attention_bias", 512), TensorCategory::Bias);
assert_eq!(classify_tensor("layer.bias", 128), TensorCategory::Bias);
}
#[test]
fn test_classify_tensor_layernorm() {
assert_eq!(
classify_tensor("layer_norm.gamma", 256),
TensorCategory::LayerNorm
);
assert_eq!(classify_tensor("ln_weight", 512), TensorCategory::LayerNorm);
assert_eq!(
classify_tensor("layernorm.beta", 256),
TensorCategory::LayerNorm
);
}
#[test]
fn test_classify_tensor_small() {
assert_eq!(classify_tensor("tiny.weight", 6), TensorCategory::Small);
assert_eq!(classify_tensor("fc.weight", 15), TensorCategory::Small);
assert_eq!(classify_tensor("scalar", 1), TensorCategory::Small);
}
#[test]
fn test_quantize_varmap_parallel_basic() {
let device = Device::Cpu;
let varmap = Arc::new(VarMap::new());
// Add various tensor types
{
let mut vars_data = varmap.data().lock().unwrap();
// Large weight (should quantize)
let weight1 = Tensor::randn(0_f32, 1.0, (128, 256), &device).unwrap();
vars_data.insert(
"fc1.weight".to_owned(),
Var::from_tensor(&weight1).unwrap(),
);
// Bias (should skip)
let bias1 = Tensor::randn(0_f32, 0.1, (256,), &device).unwrap();
vars_data.insert("fc1.bias".to_owned(), Var::from_tensor(&bias1).unwrap());
// LayerNorm (should skip)
let ln_gamma = Tensor::ones((256,), DType::F32, &device).unwrap();
vars_data.insert(
"layer_norm.gamma".to_owned(),
Var::from_tensor(&ln_gamma).unwrap(),
);
// Small tensor (should skip)
let small = Tensor::randn(0_f32, 1.0, (2, 3), &device).unwrap();
vars_data.insert("tiny.weight".to_owned(), Var::from_tensor(&small).unwrap());
// Another large weight (should quantize)
let weight2 = Tensor::randn(0_f32, 1.0, (256, 128), &device).unwrap();
vars_data.insert(
"fc2.weight".to_owned(),
Var::from_tensor(&weight2).unwrap(),
);
}
let result = quantize_varmap_parallel(&varmap, &device);
assert!(result.is_ok(), "Parallel quantization should succeed");
let quantized_map = result.unwrap();
// Should only quantize the 2 large weights
assert_eq!(quantized_map.len(), 2, "Should quantize 2 large weights");
assert!(quantized_map.contains_key("fc1.weight"));
assert!(quantized_map.contains_key("fc2.weight"));
assert!(!quantized_map.contains_key("fc1.bias"));
assert!(!quantized_map.contains_key("layer_norm.gamma"));
assert!(!quantized_map.contains_key("tiny.weight"));
}
#[test]
fn test_quantize_varmap_parallel_memory_reduction() {
let device = Device::Cpu;
let varmap = Arc::new(VarMap::new());
// Add 10 large weight tensors
{
let mut vars_data = varmap.data().lock().unwrap();
for i in 0..10 {
let name = format!("layer_{}.weight", i);
let tensor = Tensor::randn(0_f32, 1.0, (128, 128), &device).unwrap();
vars_data.insert(name, Var::from_tensor(&tensor).unwrap());
}
}
let result = quantize_varmap_parallel(&varmap, &device);
assert!(result.is_ok());
let quantized_map = result.unwrap();
assert_eq!(quantized_map.len(), 10, "Should quantize all 10 weights");
// Calculate memory reduction
let original_bytes = 10 * 128 * 128 * 4; // 10 tensors * 128*128 elements * 4 bytes
let quantized_bytes: usize = quantized_map.values().map(|q| q.memory_bytes()).sum();
let reduction_percent =
((original_bytes - quantized_bytes) as f64 / original_bytes as f64) * 100.0;
println!("Original: {} bytes", original_bytes);
println!("Quantized: {} bytes", quantized_bytes);
println!("Reduction: {:.1}%", reduction_percent);
// Should save ~75% memory
assert!(
reduction_percent > 70.0,
"Memory reduction should be >70%, got {:.1}%",
reduction_percent
);
}
#[test]
fn test_quantize_varmap_parallel_performance() {
let device = Device::Cpu;
let varmap = Arc::new(VarMap::new());
// Add 100 tensors (simulating larger model)
{
let mut vars_data = varmap.data().lock().unwrap();
for i in 0..100 {
let name = format!("layer_{}.weight", i);
let tensor = Tensor::randn(0_f32, 1.0, (64, 64), &device).unwrap();
vars_data.insert(name, Var::from_tensor(&tensor).unwrap());
}
}
// Measure quantization time
let start = std::time::Instant::now();
let result = quantize_varmap_parallel(&varmap, &device);
let elapsed = start.elapsed();
assert!(result.is_ok());
let quantized_map = result.unwrap();
assert_eq!(quantized_map.len(), 100);
// Performance target: <5s for 100 tensors (parallel should be faster)
println!("Parallel quantization time for 100 tensors: {:?}", elapsed);
assert!(
elapsed.as_secs() < 5,
"Parallel quantization should be <5s, got {:?}",
elapsed
);
}
}