Files
foxhunt/ml/tests/tft_int8_forward_integration_test.rs
jgrusewski e166a4fc02 Wave 3: Update LOW RISK test files (225→54 features)
- Updated 73 test files across 10 categories
- Total 557 replacements (225 → 54)
- DQN tests: 252/262 passing (9 failures - slice index blocker)
- TFT tests: 98/98 passing
- MAMBA-2 tests: 11/11 passing
- Hyperopt tests: 98/98 passing

Critical findings:
- Blocker: ml/src/trainers/dqn.rs:3444 hardcoded slice indices
- Architecture mismatch: extract_current_features() vs extract_current_features_v2()

Wave 3 Agent breakdown:
- Agent 1: DQN test files (12 files)
- Agent 2: PPO test files (2 files)
- Agent 3: TFT test files (6 files)
- Agent 4: MAMBA-2 test files (2 files)
- Agent 5: Feature extraction tests (3 files)
- Agent 6: Integration test files (9 files)
- Agent 7: Data loader test files (3 files)
- Agent 8: Hyperopt test files (1 file)
- Agent 9: Benchmark test files (9 files)
- Agent 10: Utility & misc test files (73 files)

Next: Fix slice index blocker, then Wave 4 (OFI integration 46→54)
2025-11-23 01:22:32 +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 foxhunt_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(())
}