408 lines
13 KiB
Rust
408 lines
13 KiB
Rust
use anyhow::Result;
|
|
use rtx_compress::quantization::mixed_precision::{
|
|
LayerSensitivity, MixedPrecisionOptimizer, OptimizationObjective, PrecisionConfig,
|
|
};
|
|
use rtx_tensor::{Device, Shape, Tensor};
|
|
use std::collections::HashMap;
|
|
|
|
#[test]
|
|
#[ignore = "Pre-existing Metal device randn issue"]
|
|
fn test_layer_sensitivity_analysis() -> Result<()> {
|
|
let device = Device::try_default()?;
|
|
|
|
// Mock model layers
|
|
let mut layers = HashMap::new();
|
|
layers.insert(
|
|
"embeddings".to_string(),
|
|
Tensor::randn(&[30000, 512], &device)?,
|
|
);
|
|
layers.insert(
|
|
"attention.query".to_string(),
|
|
Tensor::randn(&[512, 512], &device)?,
|
|
);
|
|
layers.insert(
|
|
"attention.key".to_string(),
|
|
Tensor::randn(&[512, 512], &device)?,
|
|
);
|
|
layers.insert(
|
|
"mlp.up_proj".to_string(),
|
|
Tensor::randn(&[512, 2048], &device)?,
|
|
);
|
|
layers.insert("layer_norm".to_string(), Tensor::randn(&[512], &device)?);
|
|
|
|
let config = PrecisionConfig {
|
|
precision_bits: vec![4, 8, 16],
|
|
sensitivity_threshold: 0.02,
|
|
performance_weight: 0.7,
|
|
quality_weight: 0.3,
|
|
};
|
|
|
|
let mut optimizer = MixedPrecisionOptimizer::new(config);
|
|
|
|
// Mock calibration data for sensitivity analysis
|
|
let calibration_data = Tensor::randn(&[100, 512], &device)?;
|
|
|
|
// Analyze layer sensitivities
|
|
let sensitivities = optimizer.analyze_sensitivity(&layers, &calibration_data)?;
|
|
|
|
// Verify sensitivity analysis results
|
|
assert_eq!(sensitivities.len(), layers.len());
|
|
|
|
// Layer norm should be most sensitive (normalization layers typically are)
|
|
let layer_norm_sensitivity = sensitivities.get("layer_norm").unwrap();
|
|
assert!(
|
|
layer_norm_sensitivity.error_impact > 0.05,
|
|
"Layer norm should have high sensitivity"
|
|
);
|
|
|
|
// Embeddings should be least sensitive (large, redundant parameters)
|
|
let embeddings_sensitivity = sensitivities.get("embeddings").unwrap();
|
|
assert!(
|
|
embeddings_sensitivity.error_impact < layer_norm_sensitivity.error_impact,
|
|
"Embeddings should be less sensitive than layer norm"
|
|
);
|
|
|
|
// Verify sensitivity ordering makes sense
|
|
let mut sorted_sensitivities: Vec<_> = sensitivities.iter().collect();
|
|
sorted_sensitivities.sort_by(|a, b| b.1.error_impact.partial_cmp(&a.1.error_impact).unwrap());
|
|
|
|
println!("Layer sensitivity ranking:");
|
|
for (layer, sensitivity) in &sorted_sensitivities {
|
|
println!(" {}: {:.4}", layer, sensitivity.error_impact);
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "Pre-existing Metal device randn issue"]
|
|
fn test_precision_search_optimization() -> Result<()> {
|
|
let device = Device::try_default()?;
|
|
|
|
let mut layers = HashMap::new();
|
|
for i in 0..12 {
|
|
// 12-layer transformer
|
|
layers.insert(
|
|
format!("layers.{}.attention", i),
|
|
Tensor::randn(&[768, 768], &device)?,
|
|
);
|
|
layers.insert(
|
|
format!("layers.{}.mlp", i),
|
|
Tensor::randn(&[768, 3072], &device)?,
|
|
);
|
|
}
|
|
|
|
let config = PrecisionConfig {
|
|
precision_bits: vec![4, 6, 8, 12, 16],
|
|
sensitivity_threshold: 0.01,
|
|
performance_weight: 0.6,
|
|
quality_weight: 0.4,
|
|
};
|
|
|
|
let mut optimizer = MixedPrecisionOptimizer::new(config);
|
|
|
|
let calibration_data = Tensor::randn(&[50, 768], &device)?;
|
|
|
|
// Set optimization objective
|
|
let objective = OptimizationObjective {
|
|
target_compression_ratio: 4.0,
|
|
max_quality_loss: 0.05,
|
|
memory_constraint_mb: Some(500),
|
|
};
|
|
|
|
// Run optimization
|
|
let optimal_config = optimizer.optimize(&layers, &calibration_data, objective)?;
|
|
|
|
// Verify optimization results
|
|
assert_eq!(optimal_config.layer_precisions.len(), layers.len());
|
|
|
|
// Calculate achieved compression ratio
|
|
let original_bits = layers.len() * 32; // All fp32
|
|
let compressed_bits: usize = optimal_config
|
|
.layer_precisions
|
|
.values()
|
|
.map(|&b| b as usize)
|
|
.sum();
|
|
let actual_ratio = original_bits as f32 / compressed_bits as f32;
|
|
|
|
assert!(
|
|
actual_ratio >= 3.5,
|
|
"Should achieve at least 3.5x compression"
|
|
);
|
|
assert!(
|
|
optimal_config.estimated_quality_loss <= 0.06,
|
|
"Quality loss should be within bounds"
|
|
);
|
|
|
|
// Verify precision assignment makes sense
|
|
// Earlier layers often need higher precision
|
|
let layer_0_precision = optimal_config
|
|
.layer_precisions
|
|
.get("layers.0.attention")
|
|
.unwrap();
|
|
let layer_11_precision = optimal_config
|
|
.layer_precisions
|
|
.get("layers.11.attention")
|
|
.unwrap();
|
|
|
|
// This is a heuristic - may not always hold, but generally true
|
|
println!(
|
|
"Layer 0 precision: {}, Layer 11 precision: {}",
|
|
layer_0_precision, layer_11_precision
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "Pre-existing Metal device randn issue"]
|
|
fn test_dynamic_precision_adjustment() -> Result<()> {
|
|
let device = Device::try_default()?;
|
|
|
|
let mut layers = HashMap::new();
|
|
layers.insert(
|
|
"critical_layer".to_string(),
|
|
Tensor::randn(&[256, 256], &device)?,
|
|
);
|
|
layers.insert(
|
|
"redundant_layer".to_string(),
|
|
Tensor::randn(&[1024, 1024], &device)?,
|
|
);
|
|
|
|
let config = PrecisionConfig {
|
|
precision_bits: vec![4, 8, 12, 16],
|
|
sensitivity_threshold: 0.015,
|
|
performance_weight: 0.5,
|
|
quality_weight: 0.5,
|
|
};
|
|
|
|
let mut optimizer = MixedPrecisionOptimizer::new(config);
|
|
|
|
// Initial calibration
|
|
let calibration_data = Tensor::randn(&[100, 256], &device)?;
|
|
let sensitivities = optimizer.analyze_sensitivity(&layers, &calibration_data)?;
|
|
|
|
// Simulate runtime performance feedback
|
|
let mut performance_feedback = HashMap::new();
|
|
performance_feedback.insert("critical_layer".to_string(), 0.95); // High accuracy needed
|
|
performance_feedback.insert("redundant_layer".to_string(), 0.85); // Lower accuracy ok
|
|
|
|
// Adjust precisions based on runtime feedback
|
|
let adjusted_config =
|
|
optimizer.adjust_precisions_runtime(&sensitivities, &performance_feedback)?;
|
|
|
|
// Critical layer should get higher precision
|
|
let critical_precision = adjusted_config
|
|
.layer_precisions
|
|
.get("critical_layer")
|
|
.unwrap();
|
|
let redundant_precision = adjusted_config
|
|
.layer_precisions
|
|
.get("redundant_layer")
|
|
.unwrap();
|
|
|
|
assert!(
|
|
critical_precision >= redundant_precision,
|
|
"Critical layer should have >= precision than redundant layer"
|
|
);
|
|
|
|
// Verify precision bounds
|
|
assert!(
|
|
*critical_precision >= 8,
|
|
"Critical layer needs at least 8-bit"
|
|
);
|
|
assert!(*redundant_precision >= 4, "All layers need at least 4-bit");
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "Pre-existing Metal device randn issue"]
|
|
fn test_quantization_aware_training_simulation() -> Result<()> {
|
|
let device = Device::try_default()?;
|
|
|
|
let mut layers = HashMap::new();
|
|
layers.insert("layer1".to_string(), Tensor::randn(&[128, 256], &device)?);
|
|
layers.insert("layer2".to_string(), Tensor::randn(&[256, 512], &device)?);
|
|
|
|
let config = PrecisionConfig {
|
|
precision_bits: vec![6, 8, 12],
|
|
sensitivity_threshold: 0.02,
|
|
performance_weight: 0.4,
|
|
quality_weight: 0.6,
|
|
};
|
|
|
|
let mut optimizer = MixedPrecisionOptimizer::new(config);
|
|
optimizer.enable_quantization_aware_mode(true);
|
|
|
|
// Simulate training iterations with precision updates
|
|
let mut current_precisions = HashMap::new();
|
|
current_precisions.insert("layer1".to_string(), 8);
|
|
current_precisions.insert("layer2".to_string(), 8);
|
|
|
|
for iteration in 0..10 {
|
|
// Generate training batch
|
|
let batch_data = Tensor::randn(&[32, 128], &device)?;
|
|
|
|
// Simulate forward pass with current precisions
|
|
let layer1_quantized =
|
|
optimizer.simulate_quantization(&layers["layer1"], current_precisions["layer1"])?;
|
|
|
|
// Compute simulated loss (simplified)
|
|
let two = Tensor::from_data(vec![2.0], vec![1], &device)?;
|
|
let loss = layer1_quantized
|
|
.pow(&two)?
|
|
.mean(&[], false)?
|
|
.to_scalar::<f32>()?;
|
|
|
|
// Update precisions based on loss gradients
|
|
if iteration > 5 && loss > 0.1 {
|
|
// Increase precision if loss is high
|
|
let new_precision = (current_precisions["layer1"] + 2).min(12);
|
|
current_precisions.insert("layer1".to_string(), new_precision);
|
|
}
|
|
}
|
|
|
|
// Verify precision was adjusted during "training"
|
|
let final_layer1_precision = current_precisions["layer1"];
|
|
println!("Final layer1 precision: {}", final_layer1_precision);
|
|
|
|
// Should either stay at 8 or increase based on loss feedback
|
|
assert!(final_layer1_precision >= 8 && final_layer1_precision <= 12);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "Pre-existing Metal device randn issue"]
|
|
fn test_hardware_aware_precision_mapping() -> Result<()> {
|
|
let device = Device::try_default()?;
|
|
|
|
let mut layers = HashMap::new();
|
|
for i in 0..8 {
|
|
layers.insert(
|
|
format!("conv_{}", i),
|
|
Tensor::randn(&[64, 128, 3, 3], &device)?,
|
|
);
|
|
}
|
|
|
|
let config = PrecisionConfig {
|
|
precision_bits: vec![4, 8, 16],
|
|
sensitivity_threshold: 0.01,
|
|
performance_weight: 0.8, // Prioritize performance
|
|
quality_weight: 0.2,
|
|
};
|
|
|
|
let mut optimizer = MixedPrecisionOptimizer::new(config);
|
|
|
|
// Configure hardware constraints
|
|
optimizer.set_hardware_constraints(&[
|
|
("int4_ops_per_sec", 1_000_000.0),
|
|
("int8_ops_per_sec", 500000.0),
|
|
("fp16_ops_per_sec", 200000.0),
|
|
])?;
|
|
|
|
let calibration_data = Tensor::randn(&[100, 64], &device)?;
|
|
|
|
let objective = OptimizationObjective {
|
|
target_compression_ratio: 3.0,
|
|
max_quality_loss: 0.08,
|
|
memory_constraint_mb: None,
|
|
};
|
|
|
|
let optimal_config = optimizer.optimize(&layers, &calibration_data, objective)?;
|
|
|
|
// Verify hardware-aware decisions
|
|
let precision_counts = count_precision_usage(&optimal_config.layer_precisions);
|
|
|
|
// Should prefer 4-bit (fastest) when quality allows
|
|
assert!(
|
|
precision_counts.get(&4).unwrap_or(&0) > &0,
|
|
"Should use some 4-bit quantization for performance"
|
|
);
|
|
|
|
// Calculate theoretical throughput
|
|
let estimated_ops_per_sec = calculate_throughput(&optimal_config.layer_precisions);
|
|
println!("Estimated throughput: {:.0} ops/sec", estimated_ops_per_sec);
|
|
|
|
assert!(
|
|
estimated_ops_per_sec > 300000.0,
|
|
"Should achieve reasonable throughput with hardware constraints"
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "Pre-existing Metal device randn issue"]
|
|
fn test_precision_config_serialization() -> Result<()> {
|
|
let device = Device::try_default()?;
|
|
|
|
let mut layers = HashMap::new();
|
|
layers.insert(
|
|
"test_layer".to_string(),
|
|
Tensor::randn(&[64, 128], &device)?,
|
|
);
|
|
|
|
let config = PrecisionConfig {
|
|
precision_bits: vec![4, 8],
|
|
sensitivity_threshold: 0.03,
|
|
performance_weight: 0.6,
|
|
quality_weight: 0.4,
|
|
};
|
|
|
|
let mut optimizer = MixedPrecisionOptimizer::new(config);
|
|
|
|
let calibration_data = Tensor::randn(&[50, 64], &device)?;
|
|
let sensitivities = optimizer.analyze_sensitivity(&layers, &calibration_data)?;
|
|
|
|
// Serialize sensitivity analysis results
|
|
let serialized = optimizer.serialize_sensitivities(&sensitivities)?;
|
|
|
|
// Create new optimizer and deserialize
|
|
let mut optimizer2 = MixedPrecisionOptimizer::new(PrecisionConfig {
|
|
precision_bits: vec![4, 8],
|
|
sensitivity_threshold: 0.03,
|
|
performance_weight: 0.6,
|
|
quality_weight: 0.4,
|
|
});
|
|
|
|
let deserialized_sensitivities = optimizer2.deserialize_sensitivities(&serialized)?;
|
|
|
|
// Verify serialization preserved data
|
|
assert_eq!(deserialized_sensitivities.len(), sensitivities.len());
|
|
|
|
for (layer_name, original_sens) in &sensitivities {
|
|
let deserialized_sens = deserialized_sensitivities.get(layer_name).unwrap();
|
|
let error_diff = (original_sens.error_impact - deserialized_sens.error_impact).abs();
|
|
assert!(
|
|
error_diff < 1e-6,
|
|
"Sensitivity data should be preserved exactly"
|
|
);
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
fn count_precision_usage(precisions: &HashMap<String, u8>) -> HashMap<u8, usize> {
|
|
let mut counts = HashMap::new();
|
|
for precision in precisions.values() {
|
|
*counts.entry(*precision).or_insert(0) += 1;
|
|
}
|
|
counts
|
|
}
|
|
|
|
fn calculate_throughput(precisions: &HashMap<String, u8>) -> f64 {
|
|
// Simplified throughput calculation based on precision
|
|
precisions
|
|
.values()
|
|
.map(|&bits| match bits {
|
|
4 => 1_000_000.0,
|
|
8 => 500000.0,
|
|
16 => 200000.0,
|
|
_ => 100000.0,
|
|
})
|
|
.sum::<f64>()
|
|
/ precisions.len() as f64
|
|
}
|