//! Standalone test for GAE module //! Run with: cargo test --test gae_standalone_test use ml::dqn::{GAECalculator, GAEConfig}; #[test] fn test_gae_basic_functionality() { let gae = GAECalculator::new(0.99, 0.95); let rewards = vec![1.0, 2.0, 3.0]; let values = vec![0.5, 0.6, 0.7]; let dones = vec![false, false, true]; let returns = gae.compute_returns(&rewards, &values, &dones); assert_eq!(returns.len(), 3); for r in &returns { assert!(r.is_finite()); } println!("✓ GAE basic test passed: returns = {:?}", returns); } #[test] fn test_gae_from_config() { let config = GAEConfig { gamma: 0.98, lambda: 0.9, }; let gae = GAECalculator::from_config(&config); assert_eq!(gae.gamma(), 0.98); assert_eq!(gae.lambda(), 0.9); println!("✓ GAE config test passed"); } #[test] fn test_gae_advantages_separate() { let gae = GAECalculator::new(0.99, 0.95); let rewards = vec![1.0, 2.0]; let values = vec![0.5, 0.6]; let dones = vec![false, false]; let advantages = gae.compute_advantages(&rewards, &values, &dones); let returns = gae.compute_returns(&rewards, &values, &dones); assert_eq!(advantages.len(), returns.len()); // Verify: returns = advantages + values for i in 0..advantages.len() { assert!((returns[i] - (advantages[i] + values[i])).abs() < 1e-6); } println!("✓ GAE advantages test passed"); } #[test] fn test_gae_lambda_zero_equals_td() { // Lambda = 0 should give TD(0) returns let gae = GAECalculator::new(0.99, 0.0); let rewards = vec![1.0, 2.0, 3.0]; let values = vec![0.5, 0.6, 0.7]; let dones = vec![false, false, false]; let returns = gae.compute_returns(&rewards, &values, &dones); // With λ=0, GAE reduces to TD(0): A_t = δ_t assert_eq!(returns.len(), 3); assert!((returns[2] - 3.0).abs() < 1e-6); assert!((returns[1] - 2.693).abs() < 1e-6); assert!((returns[0] - 1.594).abs() < 1e-6); println!("✓ GAE lambda=0 test passed"); } #[test] fn test_gae_episode_boundary() { let gae = GAECalculator::new(0.99, 0.95); let rewards = vec![1.0, 2.0, 3.0]; let values = vec![0.5, 0.6, 0.7]; let dones = vec![false, true, false]; // Episode ends at step 1 let returns = gae.compute_returns(&rewards, &values, &dones); assert_eq!(returns.len(), 3); // All returns should be finite for r in &returns { assert!(r.is_finite()); } println!("✓ GAE episode boundary test passed: returns = {:?}", returns); }