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>
961 lines
34 KiB
Rust
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
|
|
);
|
|
}
|
|
}
|