//! 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, // 30 - 5 - 10 = 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: 225, // Full Wave C+D features 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: 210, ..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(()) }