/// Integration test for QuantizedTFT forward() implementation /// /// Validates end-to-end forward pass with all 6 sub-methods integrated use candle_core::{DType, Device, Tensor}; use ml::memory_optimization::quantization::{QuantizationConfig, QuantizationType, Quantizer}; use ml::tft::{QuantizedTemporalFusionTransformer, TFTConfig}; use ml::MLError; use std::collections::HashMap; #[test] fn test_forward_pass_basic() -> Result<(), MLError> { // Test configuration let config = TFTConfig { input_dim: 54, hidden_dim: 256, num_heads: 8, num_layers: 4, prediction_horizon: 10, sequence_length: 60, num_quantiles: 3, num_static_features: 20, num_known_features: 10, num_unknown_features: 195, learning_rate: 0.001, batch_size: 32, dropout_rate: 0.1, l2_regularization: 0.0001, use_flash_attention: false, mixed_precision: false, memory_efficient: true, max_inference_latency_us: 3200, target_throughput_pps: 10_000, }; let device = Device::Cpu; let mut model = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())?; // Create input tensors let batch_size = 2; // Static features: [batch, num_static_features=20] let static_features = Tensor::randn(0f32, 1.0, (batch_size, config.num_static_features), &device)?; // Historical features: [batch, seq_len=60, num_unknown_features=195] let historical_features = Tensor::randn( 0f32, 1.0, ( batch_size, config.sequence_length, config.num_unknown_features, ), &device, )?; // Future features: [batch, horizon=10, num_known_features=10] let future_features = Tensor::randn( 0f32, 1.0, ( batch_size, config.prediction_horizon, config.num_known_features, ), &device, )?; // Initialize attention weights (required for forward pass) let hidden_dim = config.hidden_dim; let q_weight = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), &device)?; let k_weight = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), &device)?; let v_weight = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), &device)?; let o_weight = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), &device)?; let mut quantizer = Quantizer::new( QuantizationConfig { quant_type: QuantizationType::Int8, per_channel: false, symmetric: true, calibration_samples: None, }, device.clone(), ); let q_weight_int8 = quantizer.quantize_tensor(&q_weight, "q_weight")?; let k_weight_int8 = quantizer.quantize_tensor(&k_weight, "k_weight")?; let v_weight_int8 = quantizer.quantize_tensor(&v_weight, "v_weight")?; let o_weight_int8 = quantizer.quantize_tensor(&o_weight, "o_weight")?; model.initialize_attention_weights(q_weight_int8, k_weight_int8, v_weight_int8, o_weight_int8); // Initialize static VSN weights let mut static_vsn_weights = HashMap::new(); let vsn_weight = Tensor::randn(0f32, 0.1, (hidden_dim, config.num_static_features), &device)?; let vsn_weight_int8 = quantizer.quantize_tensor(&vsn_weight, "static_vsn")?; static_vsn_weights.insert("static_vsn".to_string(), vsn_weight_int8); model.initialize_static_vsn_weights(static_vsn_weights); // Run forward pass let output = model.forward(&static_features, &historical_features, &future_features)?; // Validate output shape: [batch=2, horizon=10, quantiles=3] assert_eq!( output.dims(), &[batch_size, config.prediction_horizon, config.num_quantiles], "Output shape mismatch" ); // Validate no NaN/Inf values let output_data = output.flatten_all()?.to_vec1::()?; assert!( output_data.iter().all(|x| x.is_finite()), "Output contains NaN or Inf values" ); println!("✅ Forward pass test passed!"); println!(" Output shape: {:?}", output.dims()); println!( " Output range: [{:.4}, {:.4}]", output_data.iter().fold(f32::INFINITY, |a, &b| a.min(b)), output_data.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b)) ); Ok(()) } #[test] fn test_forward_pass_with_device_mismatch() { let config = TFTConfig::default(); let device = Device::Cpu; let mut model = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone()) .unwrap(); let batch_size = 2; // Create inputs on correct device let static_features = Tensor::zeros( (batch_size, config.num_static_features), DType::F32, &device, ) .unwrap(); let historical_features = Tensor::zeros( ( batch_size, config.sequence_length, config.num_unknown_features, ), DType::F32, &device, ) .unwrap(); let future_features = Tensor::zeros( ( batch_size, config.prediction_horizon, config.num_known_features, ), DType::F32, &device, ) .unwrap(); // This should work (all on same device) let result = model.forward(&static_features, &historical_features, &future_features); // Should succeed even without weights initialized (falls back to zeros) assert!( result.is_ok(), "Forward pass should succeed with fallback behavior" ); } #[test] fn test_forward_pass_validates_dimensions() { let config = TFTConfig::default(); let device = Device::Cpu; let mut model = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone()) .unwrap(); let batch_size = 2; // Test 1: Wrong static features dimensions let wrong_static = Tensor::zeros((batch_size, 999), DType::F32, &device).unwrap(); let hist = Tensor::zeros( ( batch_size, config.sequence_length, config.num_unknown_features, ), DType::F32, &device, ) .unwrap(); let fut = Tensor::zeros( ( batch_size, config.prediction_horizon, config.num_known_features, ), DType::F32, &device, ) .unwrap(); let result = model.forward(&wrong_static, &hist, &fut); assert!( result.is_err(), "Should reject wrong static feature dimensions" ); // Test 2: Wrong historical features dimensions let stat = Tensor::zeros( (batch_size, config.num_static_features), DType::F32, &device, ) .unwrap(); let wrong_hist = Tensor::zeros( (batch_size, config.sequence_length, 999), DType::F32, &device, ) .unwrap(); let result = model.forward(&stat, &wrong_hist, &fut); assert!( result.is_err(), "Should reject wrong historical feature dimensions" ); // Test 3: Wrong future features dimensions let wrong_fut = Tensor::zeros( (batch_size, config.prediction_horizon, 999), DType::F32, &device, ) .unwrap(); let result = model.forward(&stat, &hist, &wrong_fut); assert!( result.is_err(), "Should reject wrong future feature dimensions" ); }