use candle_core::{Device, Tensor}; use ml::memory_optimization::quantization::Quantizer; use ml::tft::{quantized_tft::QuantizedTemporalFusionTransformer, TFTConfig}; fn main() -> Result<(), Box> { println!("Testing forward_future_decoder implementation...\n"); // Create TFT config let config = TFTConfig { input_dim: 225, hidden_dim: 256, num_heads: 8, num_known_features: 10, prediction_horizon: 10, ..Default::default() }; let device = Device::Cpu; let qtft = QuantizedTemporalFusionTransformer::new_with_device(config, device.clone())?; // Test 1: Create test future features [batch=2, horizon=10, features=10] println!("Test 1: Basic forward pass"); let batch_size = 2; let horizon = 10; let num_features = 10; let future_features = Tensor::randn(0f32, 1f32, (batch_size, horizon, num_features), &device)?; println!(" Input shape: {:?}", future_features.dims()); // Create decoder weights [hidden_dim=256, num_features=10] let weight_data: Vec = (0..256 * 10).map(|i| (i as f32 * 0.01).sin()).collect(); let weights_tensor = Tensor::from_slice(&weight_data, (256, 10), &device)?; // Create quantizer and quantize the weights let mut quantizer = ml::memory_optimization::quantization::Quantizer::new( ml::memory_optimization::quantization::QuantizationConfig { quant_type: ml::memory_optimization::quantization::QuantizationType::Int8, per_channel: false, symmetric: true, calibration_samples: None, }, device.clone(), ); let quantized_weights = quantizer.quantize_tensor(&weights_tensor, "decoder")?; // Run forward pass let output = qtft.forward_future_decoder(&future_features, &quantized_weights)?; println!(" Output shape: {:?}", output.dims()); println!(" Expected: [2, 10, 256]"); // Validate output shape assert_eq!(output.dims(), &[2, 10, 256], "Output shape mismatch!"); println!(" āœ“ Shape validation passed\n"); // Test 2: Check output is not all zeros println!("Test 2: Output non-zero validation"); let output_sum = output.sum_all()?.to_vec0::()?; println!(" Output sum: {}", output_sum); assert!(output_sum.abs() > 1e-6, "Output should not be all zeros"); println!(" āœ“ Non-zero validation passed\n"); // Test 3: Broadcasting correctness println!("Test 3: Different batch sizes"); for batch in [1, 4, 8] { let test_features = Tensor::randn(0f32, 1f32, (batch, 10, 10), &device)?; let test_output = qtft.forward_future_decoder(&test_features, &quantized_weights)?; assert_eq!(test_output.dims(), &[batch, 10, 256]); println!(" āœ“ Batch size {} works correctly", batch); } println!("\nāœ… All tests passed!"); Ok(()) }