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>
338 lines
8.7 KiB
Rust
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(())
|
|
}
|