//! Multi-Step Learning Validation Tests for Rainbow DQN //! //! Comprehensive validation of n-step return calculations including: //! - Mathematical correctness of discounted returns //! - Proper handling of terminal states //! - Batch processing efficiency //! - Integration with Rainbow DQN components //! - Edge case handling and robustness use ml_models::dqn::multi_step::*; use ml_models::error::ModelError; use candle_core::{Device, Tensor}; use proptest::prelude::*; #[cfg(test)] mod multi_step_validation_tests { use super::*; /// Test mathematical correctness of n-step return calculations #[test] fn test_multi_step_mathematical_correctness() -> Result<(), ModelError> { let config = MultiStepConfig { n_steps: 4, gamma: 0.95, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; // Create a sequence with known rewards let rewards = vec![1.0, 2.0, 3.0, 4.0, 5.0]; let transitions: Vec = rewards.iter() .enumerate() .map(|(i, &reward)| { create_multi_step_transition( vec![i as f32], 0, reward, vec![(i + 1) as f32], false, i, ) }) .collect(); for transition in &transitions { calculator.add_transition(transition.clone()); } let n_step_return = calculator.compute_n_step_return()?; // Manual calculation: 1.0 + 0.95*2.0 + 0.95^2*3.0 + 0.95^3*4.0 let expected = 1.0 + 0.95 * 2.0 + 0.95_f32.powi(2) * 3.0 + 0.95_f32.powi(3) * 4.0; let tolerance = 1e-6; assert!((n_step_return.n_step_reward - expected).abs() < tolerance, "Expected reward {}, got {}", expected, n_step_return.n_step_reward); assert_eq!(n_step_return.actual_steps, 4); assert!(!n_step_return.is_terminal); Ok(()) } /// Test early termination handling #[test] fn test_early_termination_correctness() -> Result<(), ModelError> { let config = MultiStepConfig { n_steps: 5, gamma: 0.9, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; // Create sequence that terminates early let transitions = vec![ create_multi_step_transition(vec![1.0], 0, 10.0, vec![2.0], false, 0), create_multi_step_transition(vec![2.0], 1, 20.0, vec![3.0], false, 1), create_multi_step_transition(vec![3.0], 2, 30.0, vec![0.0], true, 2), // Terminal ]; for transition in transitions { calculator.add_transition(transition); } let n_step_return = calculator.compute_n_step_return()?; // Should only accumulate 3 steps due to termination let expected = 10.0 + 0.9 * 20.0 + 0.9_f32.powi(2) * 30.0; assert!((n_step_return.n_step_reward - expected).abs() < 1e-6); assert_eq!(n_step_return.actual_steps, 3); assert!(n_step_return.is_terminal); Ok(()) } /// Test batch processing with various episode lengths #[test] fn test_batch_processing_mixed_episodes() -> Result<(), ModelError> { let config = MultiStepConfig { n_steps: 3, gamma: 0.95, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; // Create mixed episodes with different termination points let transitions = vec![ // Episode 1 (normal) create_multi_step_transition(vec![1.0], 0, 1.0, vec![2.0], false, 0), create_multi_step_transition(vec![2.0], 1, 2.0, vec![3.0], false, 1), create_multi_step_transition(vec![3.0], 2, 3.0, vec![4.0], false, 2), create_multi_step_transition(vec![4.0], 0, 4.0, vec![5.0], false, 3), // Episode 2 (early termination) create_multi_step_transition(vec![5.0], 1, 5.0, vec![6.0], false, 4), create_multi_step_transition(vec![6.0], 2, 6.0, vec![0.0], true, 5), // Terminal // Episode 3 (single step) create_multi_step_transition(vec![7.0], 0, 7.0, vec![0.0], true, 6), // Immediate terminal ]; let returns = calculator.compute_batch_returns(&transitions)?; // Should compute returns for eligible starting positions assert_eq!(returns.len(), 5); // 7 transitions - 3 steps + 1 = 5 possible returns // Verify first return (full 3-step) let expected_first = 1.0 + 0.95 * 2.0 + 0.95_f32.powi(2) * 3.0; assert!((returns[0].n_step_reward - expected_first).abs() < 1e-6); assert_eq!(returns[0].actual_steps, 3); assert!(!returns[0].is_terminal); Ok(()) } /// Test tensor conversion and target computation #[test] fn test_tensor_operations_correctness() -> Result<(), ModelError> { let device = Device::Cpu; let config = MultiStepConfig { n_steps: 2, gamma: 0.9, enabled: true, }; let calculator = MultiStepCalculator::new(config)?; // Create test returns let returns = vec![ MultiStepReturn { initial_state: vec![1.0, 2.0], action: 0, n_step_reward: 3.5, final_state: vec![3.0, 4.0], is_terminal: false, actual_steps: 2, gamma_n: 0.81, // 0.9^2 }, MultiStepReturn { initial_state: vec![5.0, 6.0], action: 1, n_step_reward: 7.2, final_state: vec![7.0, 8.0], is_terminal: true, actual_steps: 1, gamma_n: 0.9, }, ]; let batch = calculator.returns_to_tensors(&returns, &device)?; // Verify tensor shapes assert_eq!(batch.batch_size(), 2); assert_eq!(batch.states.shape().dims(), &[2, 2]); assert_eq!(batch.final_states.shape().dims(), &[2, 2]); // Test target computation let final_q_values = Tensor::new(&[[2.0, 4.0, 3.0], [1.0, 5.0, 2.0]], &device)?; let targets = batch.compute_targets(&final_q_values)?; let target_values = targets.to_vec1::()?; // First target: 3.5 + 0.81 * 4.0 * (1 - 0) = 3.5 + 3.24 = 6.74 assert!((target_values[0] - 6.74).abs() < 1e-6); // Second target: 7.2 + 0.9 * 5.0 * (1 - 1) = 7.2 + 0 = 7.2 assert!((target_values[1] - 7.2).abs() < 1e-6); Ok(()) } /// Test performance under high-frequency scenarios #[test] fn test_hft_performance_requirements() -> Result<(), ModelError> { let config = MultiStepConfig { n_steps: 3, gamma: 0.99, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; let num_transitions = 10_000; let batch_size = 1_000; // Generate large sequence of transitions let transitions: Vec = (0..num_transitions) .map(|i| { create_multi_step_transition( vec![i as f32, (i + 1) as f32], i % 3, (i % 10) as f32, vec![(i + 1) as f32, (i + 2) as f32], i % 100 == 99, // Terminal every 100 steps i, ) }) .collect(); let start_time = std::time::Instant::now(); // Process in batches let mut all_returns = Vec::new(); for chunk in transitions.chunks(batch_size) { let returns = calculator.compute_batch_returns(chunk)?; all_returns.extend(returns); } let processing_time = start_time.elapsed(); // Performance requirements for HFT assert!(processing_time.as_millis() < 100, "Processing took {} ms, exceeds 100ms limit", processing_time.as_millis()); // Verify computation correctness on large scale assert!(!all_returns.is_empty()); assert!(all_returns.len() > num_transitions - 3 * (num_transitions / batch_size)); println!("Processed {} transitions in {} ms", num_transitions, processing_time.as_millis()); Ok(()) } /// Test memory efficiency during sustained operation #[test] fn test_memory_efficiency() -> Result<(), ModelError> { let config = MultiStepConfig { n_steps: 5, gamma: 0.95, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; let iterations = 1_000; // Simulate sustained operation for i in 0..iterations { let transition = create_multi_step_transition( vec![i as f32], 0, (i % 10) as f32, vec![(i + 1) as f32], false, i, ); calculator.add_transition(transition); // Compute return when possible if calculator.can_compute_return() { let _return = calculator.compute_n_step_return()?; } // Verify buffer doesn't grow unbounded assert!(calculator.buffer_size() <= config.n_steps + 1, "Buffer size {} exceeds limit", calculator.buffer_size()); } Ok(()) } /// Test integration with different discount factors #[test] fn test_discount_factor_sensitivity() -> Result<(), ModelError> { let gamma_values = vec![0.9, 0.95, 0.99, 1.0]; let rewards = vec![1.0, 2.0, 3.0]; for gamma in gamma_values { let config = MultiStepConfig { n_steps: 3, gamma, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; for (i, &reward) in rewards.into_iter().enumerate() { let transition = create_multi_step_transition( vec![i as f32], 0, reward, vec![(i + 1) as f32], false, i, ); calculator.add_transition(transition); } let n_step_return = calculator.compute_n_step_return()?; // Manual calculation let expected = rewards[0] + gamma * rewards[1] + gamma.powi(2) * rewards[2]; assert!((n_step_return.n_step_reward - expected).abs() < 1e-6, "Gamma {} failed: expected {}, got {}", gamma, expected, n_step_return.n_step_reward); } Ok(()) } /// Test edge cases and error handling #[test] fn test_edge_cases_and_errors() { // Test empty buffer let config = MultiStepConfig::default(); let calculator = MultiStepCalculator::new(config).unwrap(); assert!(!calculator.can_compute_return()); assert!(calculator.compute_n_step_return().is_err()); // Test insufficient transitions let mut calculator = MultiStepCalculator::new(MultiStepConfig { n_steps: 5, ..Default::default() }).unwrap(); for i in 0..3 { calculator.add_transition(create_multi_step_transition( vec![i as f32], 0, 1.0, vec![(i+1) as f32], false, i )); } assert!(!calculator.can_compute_return()); // Test invalid configuration let invalid_configs = vec![ MultiStepConfig { n_steps: 0, ..Default::default() }, MultiStepConfig { gamma: 0.0, ..Default::default() }, MultiStepConfig { gamma: 1.1, ..Default::default() }, MultiStepConfig { enabled: false, ..Default::default() }, ]; for config in invalid_configs { assert!(MultiStepCalculator::new(config).is_err()); } } /// Property-based test for mathematical invariants proptest! { #[test] fn test_multi_step_invariants( n_steps in 1usize..10, gamma in 0.01f32..1.0, rewards in prop::collection::vec(0.0f32..100.0, 1..20) ) { let config = MultiStepConfig { n_steps, gamma, enabled: true, }; if let Ok(mut calculator) = MultiStepCalculator::new(config) { // Add transitions for (i, &reward) in rewards.into_iter().enumerate() { let transition = create_multi_step_transition( vec![i as f32], 0, reward, vec![(i + 1) as f32], false, i, ); calculator.add_transition(transition); } // Compute return if possible if calculator.can_compute_return() { if let Ok(n_step_return) = calculator.compute_n_step_return() { // Invariant: actual_steps should be <= n_steps prop_assert!(n_step_return.actual_steps <= n_steps); // Invariant: gamma_n should be gamma^actual_steps let expected_gamma_n = gamma.powi(n_step_return.actual_steps as i32); prop_assert!((n_step_return.gamma_n - expected_gamma_n).abs() < 1e-6); // Invariant: n_step_reward should be finite and non-negative for positive rewards prop_assert!(n_step_return.n_step_reward.is_finite()); if rewards.iter().all(|&r| r >= 0.0) { prop_assert!(n_step_return.n_step_reward >= 0.0); } } } } } #[test] fn test_discounted_return_properties( rewards in prop::collection::vec(-10.0f32..10.0, 1..10), gamma in 0.01f32..1.0 ) { let discounted = compute_discounted_return(&rewards, gamma); // Invariant: result should be finite prop_assert!(discounted.is_finite()); // Invariant: if all rewards are positive and gamma < 1, result should be positive if rewards.iter().all(|&r| r >= 0.0) && gamma < 1.0 { prop_assert!(discounted >= 0.0); } // Invariant: if gamma = 1, result should equal sum of rewards if (gamma - 1.0).abs() < 1e-6 { let sum: f32 = rewards.iter().sum(); prop_assert!((discounted - sum).abs() < 1e-6); } } } /// Test integration with Rainbow `DQN` components #[test] fn test_rainbow_dqn_integration() -> Result<(), ModelError> { let device = Device::Cpu; let config = MultiStepConfig { n_steps: 3, gamma: 0.99, enabled: true, }; let calculator = MultiStepCalculator::new(config)?; // Simulate Rainbow DQN experience let transitions = vec![ create_multi_step_transition(vec![0.1, 0.2, 0.3], 0, 1.5, vec![0.2, 0.3, 0.4], false, 0), create_multi_step_transition(vec![0.2, 0.3, 0.4], 1, 2.0, vec![0.3, 0.4, 0.5], false, 1), create_multi_step_transition(vec![0.3, 0.4, 0.5], 2, 1.0, vec![0.4, 0.5, 0.6], false, 2), ]; // Process through multi-step calculator let mut calc_clone = calculator; for transition in &transitions { calc_clone.add_transition(transition.clone()); } let n_step_return = calc_clone.compute_n_step_return()?; let returns = vec![n_step_return]; // Convert to tensors (as would be done in Rainbow DQN training) let batch = calculator.returns_to_tensors(&returns, &device)?; // Simulate Q-network output for final states let final_q_values = Tensor::new(&[[1.0, 2.5, 1.8]], &device)?; // Compute targets (as done in Rainbow DQN loss computation) let targets = batch.compute_targets(&final_q_values)?; // Verify target computation let target_value = targets.to_vec1::()?[0]; let expected_reward = 1.5 + 0.99 * 2.0 + 0.99_f32.powi(2) * 1.0; let expected_target = expected_reward + 0.99_f32.powi(3) * 2.5; // Bootstrap with max Q-value assert!((target_value - expected_target).abs() < 1e-6, "Expected target {}, got {}", expected_target, target_value); Ok(()) } /// Benchmark multi-step computation performance #[test] // Re-enabled: Performance testing now included in standard suite fn benchmark_multi_step_performance() -> Result<(), ModelError> { let config = MultiStepConfig { n_steps: 5, gamma: 0.99, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; let num_operations = 100_000; // Generate test data let transitions: Vec = (0..num_operations) .map(|i| { create_multi_step_transition( vec![(i % 100) as f32, ((i + 1) % 100) as f32], i % 3, (i % 10) as f32 + 1.0, vec![((i + 1) % 100) as f32, ((i + 2) % 100) as f32], i % 1000 == 999, // Terminal every 1000 steps i, ) }) .collect(); // Benchmark batch processing let start_time = std::time::Instant::now(); let returns = calculator.compute_batch_returns(&transitions)?; let batch_time = start_time.elapsed(); // Benchmark tensor conversion let device = Device::Cpu; let start_tensor_time = std::time::Instant::now(); let _batch = calculator.returns_to_tensors(&returns, &device)?; let tensor_time = start_tensor_time.elapsed(); println!("Multi-step Performance Benchmark:"); println!("Processed {} transitions in {} ms", num_operations, batch_time.as_millis()); println!("Tensor conversion took {} ms", tensor_time.as_millis()); println!("Throughput: {:.0} transitions/second", num_operations as f64 / batch_time.as_secs_f64()); // Performance requirements for HFT assert!(batch_time.as_millis() < 1000, "Batch processing too slow"); assert!(tensor_time.as_millis() < 100, "Tensor conversion too slow"); Ok(()) } } /// Helper functions for testing fn create_test_episode(length: usize, gamma: f32) -> Vec { (0..length) .map(|i| { create_multi_step_transition( vec![i as f32], 0, 1.0, vec![(i + 1) as f32], i == length - 1, // Last step is terminal i, ) }) .collect() } fn verify_n_step_calculation( transitions: &[MultiStepTransition], n_steps: usize, gamma: f32, expected_reward: f32, ) -> Result<(), ModelError> { let config = MultiStepConfig { n_steps, gamma, enabled: true, }; let mut calculator = MultiStepCalculator::new(config)?; for transition in transitions { calculator.add_transition(transition.clone()); } let n_step_return = calculator.compute_n_step_return()?; assert!((n_step_return.n_step_reward - expected_reward).abs() < 1e-6); Ok(()) } /// Integration test with mock Rainbow `DQN` environment #[cfg(test)] mod integration_tests { use super::*; #[test] fn test_end_to_end_rainbow_dqn_flow() -> Result<(), ModelError> { let device = Device::Cpu; let config = MultiStepConfig { n_steps: 3, gamma: 0.99, enabled: true, }; // Simulate complete Rainbow DQN training step let mut calculator = MultiStepCalculator::new(config)?; // Add episode data let episode = create_test_episode(10, config.gamma); let returns = calculator.compute_batch_returns(&episode)?; // Convert to training batch let batch = calculator.returns_to_tensors(&returns, &device)?; // Simulate Q-network forward pass let state_dim = 1; let action_dim = 3; let batch_size = batch.batch_size(); // Mock Q-values for current states let current_q_values = Tensor::rand(0.0, 1.0, (batch_size, action_dim), &device)?; // Mock Q-values for final states let final_q_values = Tensor::rand(0.0, 1.0, (batch_size, action_dim), &device)?; // Compute multi-step targets let targets = batch.compute_targets(&final_q_values)?; // Verify shapes and ranges assert_eq!(targets.shape().dims(), &[batch_size]); let target_values = targets.to_vec1::()?; assert!(target_values.iter().all(|&v| v.is_finite())); println!("Successfully processed end-to-end Rainbow DQN flow with {} returns", returns.len()); Ok(()) } }