Files
foxhunt/crates/ml/tests/test_tft_weight_cache.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

229 lines
8.0 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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"
);
}