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>
229 lines
8.0 KiB
Rust
229 lines
8.0 KiB
Rust
//! Tests for TFT weight caching functionality
|
||
//!
|
||
//! This test validates:
|
||
//! 1. Cache enable/disable functionality
|
||
//! 2. Automatic cache building on first use
|
||
//! 3. Cache invalidation on weight updates
|
||
//! 4. Cache statistics reporting
|
||
//! 5. Performance improvement with caching
|
||
|
||
use candle_core::{Device, Tensor};
|
||
use ml::memory_optimization::quantization::{QuantizationConfig, QuantizationType, Quantizer};
|
||
use ml::tft::{QuantizedTemporalFusionTransformer, TFTConfig};
|
||
|
||
#[test]
|
||
fn test_cache_enable_disable() {
|
||
let config = TFTConfig {
|
||
hidden_dim: 128,
|
||
..Default::default()
|
||
};
|
||
|
||
let mut model = QuantizedTemporalFusionTransformer::new(config).unwrap();
|
||
|
||
// Initially cache should be disabled
|
||
let (enabled, built, memory) = model.cache_stats();
|
||
assert!(!enabled, "Cache should be disabled by default");
|
||
assert!(!built, "Cache should not be built initially");
|
||
assert_eq!(memory, 0, "Memory usage should be 0 when cache not built");
|
||
|
||
// Enable cache
|
||
model.enable_cache();
|
||
let (enabled, built, memory) = model.cache_stats();
|
||
assert!(enabled, "Cache should be enabled after enable_cache()");
|
||
assert!(!built, "Cache should not be built until first use");
|
||
assert_eq!(memory, 0, "Memory usage should still be 0 before building");
|
||
|
||
// Disable cache
|
||
model.disable_cache();
|
||
let (enabled, built, memory) = model.cache_stats();
|
||
assert!(!enabled, "Cache should be disabled after disable_cache()");
|
||
assert!(!built, "Cache should be cleared on disable");
|
||
assert_eq!(memory, 0, "Memory usage should be 0 after disable");
|
||
}
|
||
|
||
#[test]
|
||
fn test_cache_invalidation_on_weight_update() {
|
||
let config = TFTConfig {
|
||
hidden_dim: 128,
|
||
..Default::default()
|
||
};
|
||
|
||
let device = Device::Cpu;
|
||
let mut model =
|
||
QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())
|
||
.unwrap();
|
||
|
||
// Create dummy quantized weights
|
||
let quant_config = QuantizationConfig {
|
||
quant_type: QuantizationType::Int8,
|
||
per_channel: false,
|
||
symmetric: true,
|
||
calibration_samples: None,
|
||
};
|
||
let mut quantizer = Quantizer::new(quant_config, device.clone());
|
||
|
||
// Create FP32 weights and quantize them
|
||
let weight_shape = (config.hidden_dim, config.hidden_dim);
|
||
let q_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
let k_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
let v_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
let o_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
|
||
let q_weight = quantizer.quantize_tensor(&q_weight_fp32, "q").unwrap();
|
||
let k_weight = quantizer.quantize_tensor(&k_weight_fp32, "k").unwrap();
|
||
let v_weight = quantizer.quantize_tensor(&v_weight_fp32, "v").unwrap();
|
||
let o_weight = quantizer.quantize_tensor(&o_weight_fp32, "o").unwrap();
|
||
|
||
// Enable cache and initialize weights
|
||
model.enable_cache();
|
||
model.initialize_attention_weights(q_weight, k_weight, v_weight, o_weight);
|
||
|
||
// Cache should be invalidated after weight initialization
|
||
let (enabled, built, _) = model.cache_stats();
|
||
assert!(enabled, "Cache should still be enabled");
|
||
assert!(!built, "Cache should be invalidated after weight update");
|
||
}
|
||
|
||
#[test]
|
||
fn test_cache_automatic_build_on_forward() {
|
||
let config = TFTConfig {
|
||
hidden_dim: 128,
|
||
sequence_length: 10,
|
||
..Default::default()
|
||
};
|
||
|
||
let device = Device::Cpu;
|
||
let mut model =
|
||
QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())
|
||
.unwrap();
|
||
|
||
// Create dummy quantized weights
|
||
let quant_config = QuantizationConfig {
|
||
quant_type: QuantizationType::Int8,
|
||
per_channel: false,
|
||
symmetric: true,
|
||
calibration_samples: None,
|
||
};
|
||
let mut quantizer = Quantizer::new(quant_config, device.clone());
|
||
|
||
let weight_shape = (config.hidden_dim, config.hidden_dim);
|
||
let q_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
let k_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
let v_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
let o_weight_fp32 = Tensor::zeros(weight_shape, candle_core::DType::F32, &device).unwrap();
|
||
|
||
let q_weight = quantizer.quantize_tensor(&q_weight_fp32, "q").unwrap();
|
||
let k_weight = quantizer.quantize_tensor(&k_weight_fp32, "k").unwrap();
|
||
let v_weight = quantizer.quantize_tensor(&v_weight_fp32, "v").unwrap();
|
||
let o_weight = quantizer.quantize_tensor(&o_weight_fp32, "o").unwrap();
|
||
|
||
model.enable_cache();
|
||
model.initialize_attention_weights(q_weight, k_weight, v_weight, o_weight);
|
||
|
||
// Create dummy input for forward pass
|
||
let batch_size = 2;
|
||
let input = Tensor::zeros(
|
||
(batch_size, config.sequence_length, config.hidden_dim),
|
||
candle_core::DType::F32,
|
||
&device,
|
||
)
|
||
.unwrap();
|
||
|
||
// First forward pass should build cache
|
||
let (_, built_before, _) = model.cache_stats();
|
||
assert!(
|
||
!built_before,
|
||
"Cache should not be built before first forward"
|
||
);
|
||
|
||
let _output = model.forward_attention_example(&input).unwrap();
|
||
|
||
let (enabled, built_after, memory) = model.cache_stats();
|
||
assert!(enabled, "Cache should still be enabled");
|
||
assert!(built_after, "Cache should be built after forward pass");
|
||
assert!(memory > 0, "Cache memory should be non-zero after building");
|
||
|
||
// Calculate expected memory: 4 weights × (128 × 128) × 4 bytes (FP32)
|
||
let expected_memory = 4 * config.hidden_dim * config.hidden_dim * 4;
|
||
assert_eq!(
|
||
memory, expected_memory,
|
||
"Cache memory should match expected size"
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn test_memory_usage_with_cache() {
|
||
let config = TFTConfig {
|
||
hidden_dim: 256,
|
||
..Default::default()
|
||
};
|
||
|
||
let mut model = QuantizedTemporalFusionTransformer::new(config.clone()).unwrap();
|
||
|
||
// Base memory without cache
|
||
let base_memory = model.memory_usage_bytes();
|
||
assert_eq!(
|
||
base_memory,
|
||
125 * 1024 * 1024,
|
||
"Base memory should be 125MB"
|
||
);
|
||
|
||
// Enable cache (but don't build it yet)
|
||
model.enable_cache();
|
||
let memory_cache_enabled = model.memory_usage_bytes();
|
||
assert_eq!(
|
||
memory_cache_enabled, base_memory,
|
||
"Memory should not change when cache enabled but not built"
|
||
);
|
||
|
||
// Manually test cache stats to simulate built cache
|
||
let (_, _, cache_size) = model.cache_stats();
|
||
|
||
// When cache is built, memory should increase
|
||
// Expected cache size: 4 weights × (256 × 256) × 4 bytes = 1,048,576 bytes (~1MB)
|
||
let expected_cache_size = 4 * 256 * 256 * 4;
|
||
|
||
// If cache were built, memory would be base + cache_size
|
||
// Since cache is not actually built yet, just verify the calculation
|
||
assert_eq!(cache_size, 0, "Cache size should be 0 when not built");
|
||
|
||
// Verify expected cache size calculation
|
||
assert_eq!(
|
||
expected_cache_size, 1_048_576,
|
||
"Expected cache size should be ~1MB for 256 hidden_dim"
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn test_cache_stats_accuracy() {
|
||
let config = TFTConfig {
|
||
hidden_dim: 64,
|
||
..Default::default()
|
||
};
|
||
|
||
let mut model = QuantizedTemporalFusionTransformer::new(config.clone()).unwrap();
|
||
|
||
// Test disabled state
|
||
let (enabled, built, memory) = model.cache_stats();
|
||
assert!(
|
||
!enabled && !built && memory == 0,
|
||
"Initial state should be disabled, not built, zero memory"
|
||
);
|
||
|
||
// Test enabled but not built state
|
||
model.enable_cache();
|
||
let (enabled, built, memory) = model.cache_stats();
|
||
assert!(
|
||
enabled && !built && memory == 0,
|
||
"Enabled state should be enabled, not built, zero memory"
|
||
);
|
||
|
||
// Test disabled after enable
|
||
model.disable_cache();
|
||
let (enabled, built, memory) = model.cache_stats();
|
||
assert!(
|
||
!enabled && !built && memory == 0,
|
||
"Disabled state should clear everything"
|
||
);
|
||
}
|