Files
rustytorch/crates/training/rtx-transformers/src/layers/retnet_tests.rs
T
2026-03-04 00:08:42 +00:00

377 lines
14 KiB
Rust

//! 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> {
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::<f32>()?;
let output_sum = output.sum(None)?.to_scalar::<f32>()?;
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 = (&parallel_output - &recurrent_output)?.abs()?;
let max_diff = diff.max()?.to_scalar::<f32>()?;
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::<f32>()?;
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::<f32>()?;
assert_eq!(val, 0.0);
}
}
// Upper triangle should have decay applied
if seq_len > 1 {
let val = decay_matrix.get(&[0, 1])?.to_scalar::<f32>()?;
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::<f32>()?;
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 - &parallel_output)?.abs()?;
let max_diff = diff.max()?.to_scalar::<f32>()?;
// 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::<f32>()?;
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::<f32>()?;
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::<f32>()?;
let output_norm = output.sqr()?.sum(None)?.sqrt()?.to_scalar::<f32>()?;
// 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<Tensor> {
// 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)
}