377 lines
14 KiB
Rust
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 = (¶llel_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 - ¶llel_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)
|
|
} |