//! Benchmark: Weight Caching for QuantizedTFT //! //! Compares inference performance with and without weight caching: //! - Cached mode: Dequantize once, reuse FP32 weights (4x memory, 2-3x faster) //! - Non-cached mode: Dequantize on every forward pass (saves memory, slower) //! //! Expected results: //! - Speed improvement: 2-3x faster with caching //! - Memory increase: 4x (INT8 โ†’ FP32) but still Result<(), MLError> { println!("๐Ÿ”ฌ TFT Weight Caching Benchmark"); println!("================================\n"); let device = Device::cuda_if_available(0).unwrap_or(Device::Cpu); println!("๐Ÿ“ Device: {:?}", device); // Create TFT config let mut config = TFTConfig { input_dim: 225, hidden_dim: 256, num_heads: 8, num_layers: 4, prediction_horizon: 10, sequence_length: 60, num_quantiles: 3, num_static_features: 5, num_known_features: 10, num_unknown_features: 210, 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, cache_dequantized_weights: true, // Will toggle this max_inference_latency_us: 3200, target_throughput_pps: 10_000, }; // Create test input let input = Tensor::randn(0f32, 1.0, (BATCH_SIZE, SEQ_LEN, config.hidden_dim), &device)?; println!("\n๐Ÿ“Š Test Configuration:"); println!(" Batch size: {}", BATCH_SIZE); println!(" Sequence length: {}", SEQ_LEN); println!(" Hidden dim: {}", config.hidden_dim); println!(" Warmup iterations: {}", NUM_WARMUP_ITERATIONS); println!(" Benchmark iterations: {}", NUM_BENCHMARK_ITERATIONS); // ======================================================================== // BENCHMARK 1: WITH CACHING (FAST PATH) // ======================================================================== println!("\n๐Ÿš€ Benchmark 1: WITH Weight Caching (Fast Path)"); println!("================================================"); config.cache_dequantized_weights = true; let mut model_cached = create_and_initialize_model(&config, &device)?; // Warmup println!(" ๐Ÿ”ฅ Warming up ({} iterations)...", NUM_WARMUP_ITERATIONS); for _ in 0..NUM_WARMUP_ITERATIONS { let _ = model_cached.forward_temporal_attention(&input, false)?; } // Benchmark println!( " โฑ๏ธ Benchmarking ({} iterations)...", NUM_BENCHMARK_ITERATIONS ); let start = Instant::now(); for _ in 0..NUM_BENCHMARK_ITERATIONS { let _ = model_cached.forward_temporal_attention(&input, false)?; } let cached_duration = start.elapsed(); let cached_avg_us = cached_duration.as_micros() / NUM_BENCHMARK_ITERATIONS as u128; println!(" โœ… Results:"); println!( " Total time: {:.2}ms", cached_duration.as_secs_f64() * 1000.0 ); println!(" Average per iteration: {}ยตs", cached_avg_us); // Memory profiling (cached) let memory_cached = estimate_model_memory(&model_cached); println!( " ๐Ÿ’พ Estimated memory: {:.2}MB", memory_cached / 1024.0 / 1024.0 ); // ======================================================================== // BENCHMARK 2: WITHOUT CACHING (SLOW PATH) // ======================================================================== println!("\n๐ŸŒ Benchmark 2: WITHOUT Weight Caching (Slow Path)"); println!("=================================================="); config.cache_dequantized_weights = false; let mut model_uncached = create_and_initialize_model(&config, &device)?; // Warmup println!(" ๐Ÿ”ฅ Warming up ({} iterations)...", NUM_WARMUP_ITERATIONS); for _ in 0..NUM_WARMUP_ITERATIONS { let _ = model_uncached.forward_temporal_attention(&input, false)?; } // Benchmark println!( " โฑ๏ธ Benchmarking ({} iterations)...", NUM_BENCHMARK_ITERATIONS ); let start = Instant::now(); for _ in 0..NUM_BENCHMARK_ITERATIONS { let _ = model_uncached.forward_temporal_attention(&input, false)?; } let uncached_duration = start.elapsed(); let uncached_avg_us = uncached_duration.as_micros() / NUM_BENCHMARK_ITERATIONS as u128; println!(" โœ… Results:"); println!( " Total time: {:.2}ms", uncached_duration.as_secs_f64() * 1000.0 ); println!(" Average per iteration: {}ยตs", uncached_avg_us); // Memory profiling (uncached) let memory_uncached = estimate_model_memory(&model_uncached); println!( " ๐Ÿ’พ Estimated memory: {:.2}MB", memory_uncached / 1024.0 / 1024.0 ); // ======================================================================== // SUMMARY // ======================================================================== println!("\n๐Ÿ“ˆ Performance Summary"); println!("====================="); let speedup = uncached_avg_us as f64 / cached_avg_us as f64; let memory_increase = (memory_cached as f64 / memory_uncached as f64) - 1.0; println!(" ๐Ÿ† Speed improvement: {:.2}x faster", speedup); println!( " ๐Ÿ’พ Memory increase: {:.1}% (+{:.2}MB)", memory_increase * 100.0, (memory_cached - memory_uncached) as f64 / 1024.0 / 1024.0 ); println!("\n Cached mode:"); println!(" - Latency: {}ยตs", cached_avg_us); println!(" - Memory: {:.2}MB", memory_cached / 1024.0 / 1024.0); println!("\n Uncached mode:"); println!(" - Latency: {}ยตs", uncached_avg_us); println!(" - Memory: {:.2}MB", memory_uncached / 1024.0 / 1024.0); // Validation println!("\nโœ… Validation:"); if speedup >= 2.0 { println!( " โœ“ Speed improvement meets target (โ‰ฅ2.0x): {:.2}x", speedup ); } else { println!( " โš ๏ธ Speed improvement below target (โ‰ฅ2.0x): {:.2}x", speedup ); } if memory_increase <= 5.0 { println!( " โœ“ Memory increase acceptable (โ‰ค5x): {:.2}x", memory_increase + 1.0 ); } else { println!( " โš ๏ธ Memory increase too high (>5x): {:.2}x", memory_increase + 1.0 ); } Ok(()) } /// Create and initialize a quantized TFT model with random weights fn create_and_initialize_model( config: &TFTConfig, device: &Device, ) -> Result { let mut model = QuantizedTemporalFusionTransformer::new_with_device(config.clone(), device.clone())?; // Create random FP32 weights let hidden_dim = config.hidden_dim; let q_weight_fp32 = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), device)?; let k_weight_fp32 = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), device)?; let v_weight_fp32 = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), device)?; let o_weight_fp32 = Tensor::randn(0f32, 0.1, (hidden_dim, hidden_dim), device)?; // Quantize weights let quant_config = QuantizationConfig { quant_type: QuantizationType::Int8, per_channel: false, symmetric: true, calibration_samples: None, }; let mut quantizer = Quantizer::new(quant_config, device.clone()); let q_weight_int8 = quantizer.quantize_tensor(&q_weight_fp32, "q_weight")?; let k_weight_int8 = quantizer.quantize_tensor(&k_weight_fp32, "k_weight")?; let v_weight_int8 = quantizer.quantize_tensor(&v_weight_fp32, "v_weight")?; let o_weight_int8 = quantizer.quantize_tensor(&o_weight_fp32, "o_weight")?; // Initialize model model.initialize_attention_weights(q_weight_int8, k_weight_int8, v_weight_int8, o_weight_int8); Ok(model) } /// Estimate model memory usage (rough approximation) fn estimate_model_memory(model: &QuantizedTemporalFusionTransformer) -> usize { let hidden_dim = model.config.hidden_dim; // INT8 weights: 4 matrices ร— (hidden_dim ร— hidden_dim) ร— 1 byte let quantized_size = 4 * hidden_dim * hidden_dim; // FP32 cache (if enabled): 4 matrices ร— (hidden_dim ร— hidden_dim) ร— 4 bytes let cache_size = if model.config.cache_dequantized_weights { 4 * hidden_dim * hidden_dim * 4 } else { 0 }; quantized_size + cache_size }