//! Unit tests for TFT Quantized Attention //! //! Tests INT8 quantized multi-head attention for TFT model. use anyhow::Result; use candle_core::{Device, Tensor}; use ml::memory_optimization::quantization::{QuantizationConfig, QuantizationType}; use ml::tft::quantized_attention::QuantizedMultiHeadAttention; #[test] fn test_quantized_attention_creation() -> Result<()> { let device = Device::Cpu; let d_model = 128; let num_heads = 8; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 100, per_channel: true, }; let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config, &device)?; // Verify attention was created successfully assert_eq!(attention.d_model(), d_model); assert_eq!(attention.num_heads(), num_heads); Ok(()) } #[test] fn test_quantized_attention_forward() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let seq_len = 10; let d_model = 128; let num_heads = 8; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 100, per_channel: true, }; let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config, &device)?; // Create input [batch, seq_len, d_model] let query = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; // Forward pass let output = attention.forward(&query, &key, &value, None)?; // Verify output shape matches input assert_eq!(output.dims(), &[batch_size, seq_len, d_model]); Ok(()) } #[test] fn test_quantized_attention_with_mask() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let seq_len = 8; let d_model = 64; let num_heads = 4; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 50, per_channel: true, }; let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config, &device)?; let query = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; // Create causal mask (lower triangular) let mask = Tensor::tril2(seq_len, candle_core::DType::F32, &device)?; let output = attention.forward(&query, &key, &value, Some(&mask))?; assert_eq!(output.dims(), &[batch_size, seq_len, d_model]); Ok(()) } #[test] fn test_quantized_attention_head_count_variations() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let seq_len = 10; let d_model = 128; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 100, per_channel: true, }; // Test different head counts (must divide d_model evenly) for num_heads in [1, 2, 4, 8] { let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config.clone(), &device)?; let query = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let output = attention.forward(&query, &key, &value, None)?; assert_eq!(output.dim(0)?, batch_size); assert_eq!(output.dim(1)?, seq_len); assert_eq!(output.dim(2)?, d_model); } Ok(()) } #[test] fn test_quantized_attention_output_range() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let seq_len = 5; let d_model = 64; let num_heads = 4; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 50, per_channel: true, }; let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config, &device)?; // Bounded input let query = Tensor::randn(0.0f32, 0.1, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 0.1, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 0.1, (batch_size, seq_len, d_model), &device)?; let output = attention.forward(&query, &key, &value, None)?; // Output should be finite let output_max = output.abs()?.max(0)?.max(0)?.max(0)?.to_scalar::()?; assert!(output_max.is_finite(), "Output should be finite"); assert!(output_max < 100.0, "Output should not explode"); Ok(()) } #[test] fn test_quantized_attention_gradient_flow() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let seq_len = 5; let d_model = 64; let num_heads = 4; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 50, per_channel: true, }; let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config, &device)?; let query = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let output = attention.forward(&query, &key, &value, None)?; // Compute loss for gradient check let loss = output.sum_all()?; loss.backward()?; // If backward() completes, gradient flow is working Ok(()) } #[test] fn test_quantized_attention_memory_efficiency() -> Result<()> { let device = Device::Cpu; let batch_size = 4; let seq_len = 20; let d_model = 256; let num_heads = 8; let quant_config = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 100, per_channel: true, }; // Quantized attention should use less memory than FP32 let attention = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config, &device)?; let query = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let output = attention.forward(&query, &key, &value, None)?; // Verify computation completed without OOM assert!(output.dims()[0] > 0); Ok(()) } #[test] fn test_quantized_attention_per_channel_vs_per_tensor() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let seq_len = 10; let d_model = 128; let num_heads = 8; // Per-channel quantization let quant_config_pc = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 100, per_channel: true, }; // Per-tensor quantization let quant_config_pt = QuantizationConfig { quantization_type: QuantizationType::INT8, calibration_samples: 100, per_channel: false, }; let attention_pc = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config_pc, &device)?; let attention_pt = QuantizedMultiHeadAttention::new(d_model, num_heads, quant_config_pt, &device)?; let query = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let key = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let value = Tensor::randn(0.0f32, 1.0, (batch_size, seq_len, d_model), &device)?; let output_pc = attention_pc.forward(&query, &key, &value, None)?; let output_pt = attention_pt.forward(&query, &key, &value, None)?; // Both should produce valid outputs assert_eq!(output_pc.dims(), output_pt.dims()); Ok(()) }