//! Comprehensive tests for RetNet (Retentive Networks) implementation //! //! RetNet: "Retentive Network: A Successor to Transformer for Large Language Models" (Microsoft 2023) //! Tests cover retention mechanism, decay factors, group normalization, and chunkwise computation #[cfg(all(test, feature = "disabled_tests"))] mod tests { use super::super::retnet::*; use crate::{Result, TransformerError}; use rtx_tensor::{Tensor, Device, DType}; fn create_test_tensor(shape: &[usize], device: &Device) -> Result { Tensor::randn(shape, device) } #[test] fn test_retnet_config_creation() { let config = RetNetConfig::new(512, 8, 4, 0.9); assert_eq!(config.d_model, 512); assert_eq!(config.num_heads, 8); assert_eq!(config.num_layers, 4); assert_eq!(config.decay_factor, 0.9); assert_eq!(config.group_norm_groups, 32); // default assert_eq!(config.chunk_size, 64); // default } #[test] fn test_retnet_config_with_custom_params() { let config = RetNetConfig::new(512, 8, 4, 0.9) .with_group_norm_groups(16) .with_chunk_size(128) .with_epsilon(1e-6); assert_eq!(config.group_norm_groups, 16); assert_eq!(config.chunk_size, 128); assert_eq!(config.epsilon, 1e-6); } #[test] fn test_retention_mechanism_parallel_mode() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(256, 4, 1, 0.95); let retention = RetentionMechanism::new(&config, &device)?; // Test parallel computation mode let batch_size = 2; let seq_len = 10; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output = retention.forward_parallel(&input)?; // Check output shape matches input assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); // Verify output is not identical to input (transformation occurred) let input_sum = input.sum(None)?.to_scalar::()?; let output_sum = output.sum(None)?.to_scalar::()?; assert_ne!(input_sum, output_sum); Ok(()) } #[test] fn test_retention_mechanism_recurrent_mode() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(256, 4, 1, 0.95); let retention = RetentionMechanism::new(&config, &device)?; // Test recurrent computation mode (O(n) complexity) let batch_size = 2; let seq_len = 10; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let mut state = RetentionState::new(batch_size, config.d_model, &device)?; let output = retention.forward_recurrent(&input, &mut state)?; // Check output shape assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); // State should be updated after processing assert!(state.is_initialized()); Ok(()) } #[test] fn test_retention_parallel_vs_recurrent_equivalence() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(128, 2, 1, 0.9); let retention = RetentionMechanism::new(&config, &device)?; // Create identical inputs let batch_size = 1; let seq_len = 8; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; // Parallel forward pass let parallel_output = retention.forward_parallel(&input)?; // Recurrent forward pass let mut state = RetentionState::new(batch_size, config.d_model, &device)?; let recurrent_output = retention.forward_recurrent(&input, &mut state)?; // Results should be approximately equal (within numerical precision) let diff = (¶llel_output - &recurrent_output)?.abs()?; let max_diff = diff.max()?.to_scalar::()?; assert!(max_diff < 1e-4, "Parallel and recurrent modes should produce similar results"); Ok(()) } #[test] fn test_decay_factor_application() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); // Test different decay factors let decay_factors = [0.5, 0.9, 0.95, 0.99]; for &decay in &decay_factors { let config = RetNetConfig::new(64, 2, 1, decay); let retention = RetentionMechanism::new(&config, &device)?; let batch_size = 1; let seq_len = 5; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output = retention.forward_parallel(&input)?; // Higher decay factors should preserve more information from earlier positions // This is a basic sanity check - actual behavior depends on implementation details assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); } Ok(()) } #[test] fn test_positional_decay_computation() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(128, 4, 1, 0.9); // Test decay computation for different sequence lengths let seq_lens = [1, 5, 10, 20]; for seq_len in seq_lens { let decay_matrix = compute_decay_matrix(seq_len, config.decay_factor, &device)?; // Should be upper triangular (causal masking) assert_eq!(decay_matrix.dims(), &[seq_len, seq_len]); // Diagonal should be 1.0 (no decay for current position) for i in 0..seq_len { let diag_val = decay_matrix.get(&[i, i])?.to_scalar::()?; assert!((diag_val - 1.0).abs() < 1e-6); } // Lower triangle should be 0.0 (causal masking) for i in 1..seq_len { for j in 0..i { let val = decay_matrix.get(&[i, j])?.to_scalar::()?; assert_eq!(val, 0.0); } } // Upper triangle should have decay applied if seq_len > 1 { let val = decay_matrix.get(&[0, 1])?.to_scalar::()?; assert!((val - config.decay_factor).abs() < 1e-6); } } Ok(()) } #[test] fn test_group_normalization_usage() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(256, 4, 1, 0.95) .with_group_norm_groups(8); let retnet_block = RetNetBlock::new(&config, &device)?; let batch_size = 2; let seq_len = 10; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output = retnet_block.forward(&input)?; // Group normalization should be applied internally assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); // Output should be normalized (basic sanity check) let output_mean = output.mean_keepdim(&[2])?; let mean_val = output_mean.get(&[0, 0, 0])?.to_scalar::()?; assert!(mean_val.abs() < 0.1, "Group normalization should center values near zero"); Ok(()) } #[test] fn test_chunkwise_recurrent_computation() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(128, 4, 1, 0.9) .with_chunk_size(4); let retention = RetentionMechanism::new(&config, &device)?; let batch_size = 1; let seq_len = 12; // Multiple of chunk_size let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; // Process in chunks let output = retention.forward_chunkwise(&input)?; assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); // Compare with regular parallel processing let parallel_output = retention.forward_parallel(&input)?; let diff = (&output - ¶llel_output)?.abs()?; let max_diff = diff.max()?.to_scalar::()?; // Chunkwise should approximate parallel computation assert!(max_diff < 0.01, "Chunkwise computation should approximate parallel results"); Ok(()) } #[test] fn test_retnet_block_full_forward() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(256, 8, 1, 0.95); let block = RetNetBlock::new(&config, &device)?; let batch_size = 2; let seq_len = 10; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output = block.forward(&input)?; // Check output shape assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); // Should have residual connection (output != pure retention output) let retention = RetentionMechanism::new(&config, &device)?; let retention_only = retention.forward_parallel(&input)?; let diff = (&output - &retention_only)?.abs()?; let max_diff = diff.max()?.to_scalar::()?; assert!(max_diff > 1e-3, "Block should add residual connections and normalization"); Ok(()) } #[test] fn test_retnet_complexity_scaling() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(64, 2, 1, 0.9); let retention = RetentionMechanism::new(&config, &device)?; let batch_size = 1; // Test different sequence lengths to verify O(n) recurrent scaling let seq_lens = [8, 16, 32]; let mut times = Vec::new(); for &seq_len in &seq_lens { let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let mut state = RetentionState::new(batch_size, config.d_model, &device)?; let start = std::time::Instant::now(); let _output = retention.forward_recurrent(&input, &mut state)?; let duration = start.elapsed(); times.push(duration.as_nanos() as f64); } // In recurrent mode, time should scale roughly linearly with sequence length // This is a basic check - actual timing can vary due to system factors assert!(times[1] > times[0], "Processing time should increase with sequence length"); assert!(times[2] > times[1], "Processing time should increase with sequence length"); Ok(()) } #[test] fn test_retnet_state_management() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(128, 4, 1, 0.9); let retention = RetentionMechanism::new(&config, &device)?; let batch_size = 2; let seq_len = 5; // Create state let mut state = RetentionState::new(batch_size, config.d_model, &device)?; assert!(!state.is_initialized()); // Process first sequence let input1 = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output1 = retention.forward_recurrent(&input1, &mut state)?; assert!(state.is_initialized()); // Process second sequence (should use state from first) let input2 = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output2 = retention.forward_recurrent(&input2, &mut state)?; // Outputs should be different due to state persistence let diff = (&output1 - &output2)?.abs()?; let max_diff = diff.max()?.to_scalar::()?; assert!(max_diff > 1e-3, "Different inputs should produce different outputs"); // Reset state state.reset(); assert!(!state.is_initialized()); Ok(()) } #[test] fn test_retnet_multi_layer() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(128, 4, 3, 0.95); // 3 layers let retnet = RetNet::new(&config, &device)?; let batch_size = 1; let seq_len = 8; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; let output = retnet.forward(&input)?; assert_eq!(output.dims(), &[batch_size, seq_len, config.d_model]); // Multi-layer network should significantly transform input let input_norm = input.sqr()?.sum(None)?.sqrt()?.to_scalar::()?; let output_norm = output.sqr()?.sum(None)?.sqrt()?.to_scalar::()?; // Should maintain reasonable magnitude assert!(output_norm > 0.1 * input_norm); assert!(output_norm < 10.0 * input_norm); Ok(()) } #[test] fn test_retnet_gradient_flow() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); let config = RetNetConfig::new(64, 2, 2, 0.9); let retnet = RetNet::new(&config, &device)?; let batch_size = 1; let seq_len = 4; let input = create_test_tensor(&[batch_size, seq_len, config.d_model], &device)?; // Test that parameters have gradients after forward pass let output = retnet.forward(&input)?; let loss = output.sum(None)?; // This is a basic structure test - actual gradient computation // would require autograd integration assert_eq!(loss.dims(), &[]); Ok(()) } } // Helper functions for testing fn compute_decay_matrix(seq_len: usize, decay_factor: f32, device: &Device) -> Result { // Create upper triangular matrix with decay factors let mut data = vec![0.0f32; seq_len * seq_len]; for i in 0..seq_len { for j in i..seq_len { let decay_power = (j - i) as f32; data[i * seq_len + j] = decay_factor.powf(decay_power); } } Tensor::from_slice(&data, &[seq_len, seq_len], device) }