Files
foxhunt/ml/tests/tft_int8_forward_integration_test.rs
jgrusewski 27ada2ff58 fix(ml): fix test files using wrong foxhunt_ml:: crate name
Replaced foxhunt_ml:: with ml:: in 4 test files:
- dqn_full_gradient_flow_integration_test.rs
- dqn_gradient_flow_isolation_test.rs
- tft_int8_forward_integration_test.rs
- tft_int8_integration_test.rs

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-21 13:40:25 +01:00

338 lines
8.7 KiB
Rust

//! INT8 TFT Forward Pass Integration Test
//!
//! Tests complete end-to-end forward pass through QuantizedTFT
use anyhow::Result;
use candle_core::{DType, Device, Tensor};
use ml::tft::{QuantizedTemporalFusionTransformer, TFTConfig};
#[test]
fn test_quantized_tft_forward_pass_integration() -> Result<()> {
// Create TFT configuration
let config = TFTConfig {
input_dim: 30,
hidden_dim: 64,
num_heads: 4,
num_layers: 2,
prediction_horizon: 10,
sequence_length: 20,
num_quantiles: 9,
num_static_features: 5,
num_known_features: 10,
num_unknown_features: 15,
..Default::default()
};
// Create quantized TFT model
let device = Device::Cpu;
let tft = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())?;
// Create input tensors
let batch_size = 4;
let seq_len = config.sequence_length;
let horizon = config.prediction_horizon;
let static_features = Tensor::randn(
0f32,
1f32,
(batch_size, config.num_static_features),
&device,
)?;
let historical_features = Tensor::randn(
0f32,
1f32,
(batch_size, seq_len, config.num_unknown_features),
&device,
)?;
let future_features = Tensor::randn(
0f32,
1f32,
(batch_size, horizon, config.num_known_features),
&device,
)?;
// Run forward pass
let predictions = tft.forward(&static_features, &historical_features, &future_features)?;
// Verify output shape
assert_eq!(
predictions.dims(),
&[batch_size, horizon, config.num_quantiles]
);
assert_eq!(predictions.dtype(), DType::F32);
println!("✓ Forward pass completed successfully");
println!(" Output shape: {:?}", predictions.dims());
println!(
" Memory usage: {} MB",
tft.memory_usage_bytes() / (1024 * 1024)
);
Ok(())
}
#[test]
fn test_quantized_tft_input_validation() -> Result<()> {
let config = TFTConfig {
input_dim: 30,
hidden_dim: 64,
num_heads: 4,
num_layers: 2,
prediction_horizon: 10,
sequence_length: 20,
num_quantiles: 9,
num_static_features: 5,
num_known_features: 10,
num_unknown_features: 15,
..Default::default()
};
let device = Device::Cpu;
let tft = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())?;
let batch_size = 2;
// Test 1: Invalid static features dimension
{
let invalid_static = Tensor::zeros((batch_size, 10), DType::F32, &device)?; // Wrong dim
let historical = Tensor::zeros(
(
batch_size,
config.sequence_length,
config.num_unknown_features,
),
DType::F32,
&device,
)?;
let future = Tensor::zeros(
(
batch_size,
config.prediction_horizon,
config.num_known_features,
),
DType::F32,
&device,
)?;
let result = tft.forward(&invalid_static, &historical, &future);
assert!(result.is_err(), "Should reject invalid static features");
}
// Test 2: Invalid historical features dimension
{
let static_feat = Tensor::zeros(
(batch_size, config.num_static_features),
DType::F32,
&device,
)?;
let invalid_historical = Tensor::zeros((batch_size, 20, 50), DType::F32, &device)?; // Wrong dim
let future = Tensor::zeros(
(
batch_size,
config.prediction_horizon,
config.num_known_features,
),
DType::F32,
&device,
)?;
let result = tft.forward(&static_feat, &invalid_historical, &future);
assert!(result.is_err(), "Should reject invalid historical features");
}
// Test 3: Valid inputs
{
let static_feat = Tensor::zeros(
(batch_size, config.num_static_features),
DType::F32,
&device,
)?;
let historical = Tensor::zeros(
(
batch_size,
config.sequence_length,
config.num_unknown_features,
),
DType::F32,
&device,
)?;
let future = Tensor::zeros(
(
batch_size,
config.prediction_horizon,
config.num_known_features,
),
DType::F32,
&device,
)?;
let result = tft.forward(&static_feat, &historical, &future);
assert!(result.is_ok(), "Should accept valid inputs");
}
println!("✓ Input validation tests passed");
Ok(())
}
#[test]
fn test_quantized_tft_batch_consistency() -> Result<()> {
let config = TFTConfig {
input_dim: 30,
hidden_dim: 64,
num_heads: 4,
num_layers: 2,
prediction_horizon: 10,
sequence_length: 20,
num_quantiles: 9,
num_static_features: 5,
num_known_features: 10,
num_unknown_features: 15,
..Default::default()
};
let device = Device::Cpu;
let tft = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())?;
// Test different batch sizes
for batch_size in [1, 2, 4, 8] {
let static_feat = Tensor::randn(
0f32,
1f32,
(batch_size, config.num_static_features),
&device,
)?;
let historical = Tensor::randn(
0f32,
1f32,
(
batch_size,
config.sequence_length,
config.num_unknown_features,
),
&device,
)?;
let future = Tensor::randn(
0f32,
1f32,
(
batch_size,
config.prediction_horizon,
config.num_known_features,
),
&device,
)?;
let predictions = tft.forward(&static_feat, &historical, &future)?;
assert_eq!(
predictions.dims(),
&[batch_size, config.prediction_horizon, config.num_quantiles],
"Batch size {} failed",
batch_size
);
}
println!("✓ Batch consistency tests passed");
Ok(())
}
#[test]
fn test_quantized_tft_device_consistency() -> Result<()> {
let config = TFTConfig {
input_dim: 30,
hidden_dim: 64,
num_heads: 4,
num_layers: 2,
prediction_horizon: 10,
sequence_length: 20,
num_quantiles: 9,
num_static_features: 5,
num_known_features: 10,
num_unknown_features: 15,
..Default::default()
};
// Test on CPU
let device = Device::Cpu;
let tft = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())?;
let batch_size = 2;
let static_feat = Tensor::randn(
0f32,
1f32,
(batch_size, config.num_static_features),
&device,
)?;
let historical = Tensor::randn(
0f32,
1f32,
(
batch_size,
config.sequence_length,
config.num_unknown_features,
),
&device,
)?;
let future = Tensor::randn(
0f32,
1f32,
(
batch_size,
config.prediction_horizon,
config.num_known_features,
),
&device,
)?;
let predictions = tft.forward(&static_feat, &historical, &future)?;
// Verify output is on same device
assert_eq!(predictions.device(), &device);
println!("✓ Device consistency test passed");
Ok(())
}
#[test]
fn test_quantized_tft_memory_usage() -> Result<()> {
let config = TFTConfig {
input_dim: 54,
hidden_dim: 128,
num_heads: 8,
num_layers: 3,
prediction_horizon: 10,
sequence_length: 50,
num_quantiles: 9,
num_static_features: 5,
num_known_features: 10,
num_unknown_features: 39,
..Default::default()
};
let device = Device::Cpu;
let tft = QuantizedTemporalFusionTransformer::new_with_device(config, device)?;
let memory_mb = tft.memory_usage_bytes() / (1024 * 1024);
// INT8 TFT should use ~125MB (vs 500MB for FP32)
assert!(
memory_mb <= 150,
"Memory usage {} MB exceeds 150 MB target",
memory_mb
);
assert!(
memory_mb >= 100,
"Memory usage {} MB too low, expected ~125 MB",
memory_mb
);
println!("✓ Memory usage test passed: {} MB", memory_mb);
Ok(())
}