Files
rustytorch/crates/training/rtx-compress/tests/checkpoint_tests.rs
T
2026-03-04 00:08:42 +00:00

381 lines
12 KiB
Rust

use anyhow::Result;
use rtx_compress::checkpoint::{CheckpointCompressor, CheckpointFormat, CompressionConfig};
use rtx_tensor::{Device, Tensor};
use std::collections::HashMap;
#[test]
#[ignore = "Pre-existing Metal device randn issue"]
fn test_checkpoint_save_load_compression() -> Result<()> {
let device = Device::try_default()?;
// Create a mock model state dict
let mut state_dict = HashMap::new();
state_dict.insert(
"layer1.weight".to_string(),
Tensor::randn(&[512, 1024], &device)?,
);
state_dict.insert("layer1.bias".to_string(), Tensor::randn(&[512], &device)?);
state_dict.insert(
"layer2.weight".to_string(),
Tensor::randn(&[256, 512], &device)?,
);
state_dict.insert(
"embedding.weight".to_string(),
Tensor::randn(&[30000, 512], &device)?,
);
let config = CompressionConfig {
format: CheckpointFormat::Zstd,
compression_level: 6,
quantization_bits: 8,
exclude_patterns: vec!["*.bias".to_string()], // Don't quantize biases
};
let compressor = CheckpointCompressor::new(config);
// Save compressed checkpoint
let compressed_data = compressor.save(&state_dict)?;
// Calculate compression ratio
let original_size = calculate_state_dict_size(&state_dict);
let compressed_size = compressed_data.len();
let compression_ratio = original_size as f64 / compressed_size as f64;
assert!(
compression_ratio >= 3.0,
"Compression ratio should be >= 3x, got {:.2}",
compression_ratio
);
// Load and verify
let loaded_state_dict = compressor.load(&compressed_data)?;
// Verify all keys present
assert_eq!(loaded_state_dict.len(), state_dict.len());
// Check reconstruction quality
for (key, original_tensor) in &state_dict {
let loaded_tensor = loaded_state_dict.get(key).unwrap();
assert_eq!(loaded_tensor.shape(), original_tensor.shape());
let mse = (original_tensor - loaded_tensor)?
.pow_scalar(2.0)?
.mean(&[], false)?
.to_scalar::<f32>()?;
if key.ends_with(".bias") {
// Biases should be uncompressed (perfect reconstruction)
assert!(
mse < 1e-6,
"Bias {} should have perfect reconstruction",
key
);
} else {
// Weights can have some reconstruction error
assert!(mse < 0.01, "Weight {} MSE {} too high", key, mse);
}
}
Ok(())
}
#[test]
#[ignore = "Pre-existing Metal device randn issue"]
fn test_checkpoint_mixed_precision_quantization() -> Result<()> {
let device = Device::try_default()?;
let mut state_dict = HashMap::new();
// Large embedding layer (good candidate for aggressive quantization)
state_dict.insert(
"embeddings.weight".to_string(),
Tensor::randn(&[50000, 768], &device)?,
);
// Attention weights (need higher precision)
state_dict.insert(
"attention.query.weight".to_string(),
Tensor::randn(&[768, 768], &device)?,
);
// Layer norm (should stay in fp32)
state_dict.insert(
"layer_norm.weight".to_string(),
Tensor::randn(&[768], &device)?,
);
let config = CompressionConfig {
format: CheckpointFormat::Lz4,
compression_level: 1,
quantization_bits: 0, // Use mixed precision rules
exclude_patterns: vec!["*layer_norm*".to_string()],
};
let mut compressor = CheckpointCompressor::new(config);
// Configure mixed precision rules
compressor.add_quantization_rule("*embeddings*", 4); // 4-bit for embeddings
compressor.add_quantization_rule("*attention*", 8); // 8-bit for attention
let compressed_data = compressor.save(&state_dict)?;
let loaded_state_dict = compressor.load(&compressed_data)?;
// Verify precision levels applied correctly
let embeddings_error = compute_reconstruction_error(
&state_dict["embeddings.weight"],
&loaded_state_dict["embeddings.weight"],
)?;
let attention_error = compute_reconstruction_error(
&state_dict["attention.query.weight"],
&loaded_state_dict["attention.query.weight"],
)?;
let layernorm_error = compute_reconstruction_error(
&state_dict["layer_norm.weight"],
&loaded_state_dict["layer_norm.weight"],
)?;
// Embeddings (4-bit) should have higher error than attention (8-bit)
assert!(
embeddings_error > attention_error * 2.0,
"4-bit embeddings should have more error than 8-bit attention"
);
// Layer norm should be nearly perfect (excluded from quantization)
assert!(
layernorm_error < 1e-6,
"Layer norm should have perfect reconstruction"
);
Ok(())
}
#[test]
#[ignore = "Pre-existing Metal device randn issue"]
fn test_checkpoint_incremental_compression() -> Result<()> {
let device = Device::try_default()?;
let config = CompressionConfig {
format: CheckpointFormat::Zstd,
compression_level: 3,
quantization_bits: 8,
exclude_patterns: vec![],
};
let mut compressor = CheckpointCompressor::new(config);
// Initial checkpoint
let mut state_dict_v1 = HashMap::new();
state_dict_v1.insert(
"layer1.weight".to_string(),
Tensor::randn(&[256, 512], &device)?,
);
let checkpoint_v1 = compressor.save(&state_dict_v1)?;
// Updated checkpoint (simulating fine-tuning)
let mut state_dict_v2 = state_dict_v1.clone();
state_dict_v2.insert(
"layer2.weight".to_string(),
Tensor::randn(&[128, 256], &device)?,
);
// Create incremental checkpoint
let delta_checkpoint = compressor.save_delta(&state_dict_v1, &state_dict_v2)?;
// Delta should be much smaller than full checkpoint
let full_v2 = compressor.save(&state_dict_v2)?;
assert!(
delta_checkpoint.len() < full_v2.len() / 2,
"Delta checkpoint should be significantly smaller"
);
// Apply delta to reconstruct v2
let reconstructed_v2 = compressor.apply_delta(&checkpoint_v1, &delta_checkpoint)?;
let loaded_v2 = compressor.load(&reconstructed_v2)?;
// Verify reconstruction
assert_eq!(loaded_v2.len(), 2);
assert!(loaded_v2.contains_key("layer1.weight"));
assert!(loaded_v2.contains_key("layer2.weight"));
Ok(())
}
#[test]
#[ignore = "Pre-existing Metal device randn issue"]
fn test_checkpoint_format_compatibility() -> Result<()> {
let device = Device::try_default()?;
let mut state_dict = HashMap::new();
state_dict.insert(
"test.weight".to_string(),
Tensor::randn(&[100, 200], &device)?,
);
let formats = [
CheckpointFormat::Lz4,
CheckpointFormat::Zstd,
CheckpointFormat::Uncompressed,
];
for format in &formats {
let config = CompressionConfig {
format: *format,
compression_level: 1,
quantization_bits: 16, // fp16
exclude_patterns: vec![],
};
let compressor = CheckpointCompressor::new(config);
let compressed = compressor.save(&state_dict)?;
let loaded = compressor.load(&compressed)?;
// Verify format metadata
let metadata = compressor.get_metadata(&compressed)?;
assert_eq!(metadata.format, *format);
assert_eq!(metadata.quantization_bits, 16);
// Verify reconstruction
let original = &state_dict["test.weight"];
let reconstructed = &loaded["test.weight"];
let mse = (original - reconstructed)?
.pow_scalar(2.0)?
.mean(&[], false)?
.to_scalar::<f32>()?;
assert!(
mse < 0.001,
"Format {:?} reconstruction error too high",
format
);
}
Ok(())
}
#[test]
#[ignore = "Pre-existing Metal device randn issue"]
fn test_checkpoint_large_model_streaming() -> Result<()> {
let device = Device::try_default()?;
// Simulate a large model that doesn't fit in memory
let mut state_dict = HashMap::new();
for layer in 0..96 {
// 96-layer model
state_dict.insert(
format!("layers.{}.weight", layer),
Tensor::randn(&[4096, 4096], &device)?,
);
}
let config = CompressionConfig {
format: CheckpointFormat::Zstd,
compression_level: 6,
quantization_bits: 8,
exclude_patterns: vec![],
};
let compressor = CheckpointCompressor::new(config);
// Save with streaming (should handle memory efficiently)
let mut stream = compressor.create_save_stream()?;
for (key, tensor) in &state_dict {
stream.add_tensor(key, tensor)?;
}
let compressed_data = stream.finalize()?;
// Load with streaming
let mut load_stream = compressor.create_load_stream(&compressed_data)?;
let mut loaded_count = 0;
while let Some((key, tensor)) = load_stream.next_tensor()? {
assert!(state_dict.contains_key(&key));
assert_eq!(tensor.shape(), state_dict[&key].shape());
loaded_count += 1;
}
assert_eq!(loaded_count, 96, "Should load all 96 layers");
Ok(())
}
#[test]
#[ignore = "Pre-existing Metal device randn issue"]
fn test_checkpoint_compression_benchmarks() -> Result<()> {
let device = Device::try_default()?;
// Create various tensor sizes to benchmark
let test_cases = vec![
("small", Tensor::randn(&[128, 256], &device)?),
("medium", Tensor::randn(&[1024, 2048], &device)?),
("large", Tensor::randn(&[4096, 8_192], &device)?),
];
let configs = vec![
(
"lz4_8bit",
CompressionConfig {
format: CheckpointFormat::Lz4,
compression_level: 1,
quantization_bits: 8,
exclude_patterns: vec![],
},
),
(
"zstd_4bit",
CompressionConfig {
format: CheckpointFormat::Zstd,
compression_level: 6,
quantization_bits: 4,
exclude_patterns: vec![],
},
),
];
for (case_name, tensor) in &test_cases {
for (config_name, config) in &configs {
let compressor = CheckpointCompressor::new(config.clone());
let mut state_dict = HashMap::new();
state_dict.insert("weight".to_string(), tensor.clone());
let start = std::time::Instant::now();
let compressed = compressor.save(&state_dict)?;
let save_time = start.elapsed();
let start = std::time::Instant::now();
let loaded = compressor.load(&compressed)?;
let load_time = start.elapsed();
let original_size = tensor.numel() * 4; // f32
let compressed_size = compressed.len();
let compression_ratio = original_size as f64 / compressed_size as f64;
// Log benchmark results
println!(
"Benchmark {}/{}: ratio={:.2}x, save={:?}, load={:?}",
case_name, config_name, compression_ratio, save_time, load_time
);
// Basic sanity checks
assert!(compression_ratio > 1.0);
assert!(save_time.as_secs() < 10);
assert!(load_time.as_secs() < 10);
}
}
Ok(())
}
fn calculate_state_dict_size(state_dict: &HashMap<String, Tensor>) -> usize {
state_dict
.values()
.map(|t| t.numel() * t.dtype().size_bytes())
.sum()
}
fn compute_reconstruction_error(original: &Tensor, reconstructed: &Tensor) -> Result<f32> {
let mse = (original - reconstructed)?
.pow_scalar(2.0)?
.mean(&[], false)?
.to_scalar::<f32>()?;
Ok(mse)
}