//! Convergence Validation Tests for Continuous PPO //! //! Focused tests for verifying convergence properties: //! - Policy loss convergence (monotonic decrease) //! - Value loss convergence (stabilization) //! - Exploration decay (log_std decreases) //! - Reward improvement (increasing returns) //! - Gradient stability (no explosions/vanishing) mod ppo_continuous_test_helpers; use ppo_continuous_test_helpers::{create_test_config, train_continuous_ppo, TrainingMetrics}; use anyhow::Result; #[test] fn test_policy_loss_convergence() -> Result<()> { println!("\n=== Test: Policy Loss Convergence (20 epochs) ==="); let config = create_test_config(64); // Train for 20 epochs let (_ppo, metrics) = train_continuous_ppo(20, config, 10, 20)?; println!(" Policy Loss Trajectory:"); for (i, &loss) in metrics.policy_losses.iter().enumerate() { if i % 5 == 0 || i == metrics.policy_losses.len() - 1 { println!(" Epoch {:2}: {:.4}", i + 1, loss); } } // Verify monotonic decrease (with tolerance for noise) let first_loss = metrics.policy_losses[0]; let last_loss = *metrics.policy_losses.last().unwrap(); println!("\n Convergence Analysis:"); println!(" Initial loss: {:.4}", first_loss); println!(" Final loss: {:.4}", last_loss); println!(" Reduction: {:.2}%", ((first_loss - last_loss) / first_loss.abs().max(1e-8)) * 100.0); // Assert at least 50% reduction (or final loss < threshold) let reduction_ratio = (first_loss - last_loss) / first_loss.abs().max(1e-8); assert!( reduction_ratio >= 0.5 || last_loss < 1.0, "Policy loss should reduce by ≥50% (got {:.2}%) or reach <1.0 (got {:.4})", reduction_ratio * 100.0, last_loss ); // Verify no divergence (loss doesn't explode) for (i, &loss) in metrics.policy_losses.iter().enumerate() { assert!( loss.is_finite() && loss < 1000.0, "Policy loss at epoch {} ({:.4}) should be finite and reasonable", i + 1, loss ); } println!("✓ Policy loss converges successfully"); Ok(()) } #[test] fn test_value_loss_convergence() -> Result<()> { println!("\n=== Test: Value Loss Convergence (20 epochs) ==="); let config = create_test_config(64); // Train for 20 epochs let (_ppo, metrics) = train_continuous_ppo(20, config, 10, 20)?; println!(" Value Loss Trajectory:"); for (i, &loss) in metrics.value_losses.iter().enumerate() { if i % 5 == 0 || i == metrics.value_losses.len() - 1 { println!(" Epoch {:2}: {:.4}", i + 1, loss); } } // Verify convergence to stable value let window = 5; let is_stable = metrics.is_value_loss_stable(window, 0.1); println!("\n Stability Analysis:"); if metrics.value_losses.len() >= window { let recent: Vec = metrics .value_losses .iter() .rev() .take(window) .copied() .collect(); let mean = recent.iter().sum::() / recent.len() as f32; let variance = recent.iter().map(|&x| (x - mean).powi(2)).sum::() / recent.len() as f32; let std = variance.sqrt(); println!(" Last {} epochs mean: {:.4}", window, mean); println!(" Standard deviation: {:.4}", std); println!(" Stable (std < 0.1): {}", is_stable); } // Assert stability or decreasing trend let first_loss = metrics.value_losses[0]; let last_loss = *metrics.value_losses.last().unwrap(); println!(" Initial loss: {:.4}", first_loss); println!(" Final loss: {:.4}", last_loss); // Value loss should either stabilize or decrease assert!( is_stable || last_loss < first_loss * 1.1, "Value loss should stabilize or decrease (stable={}, first={:.4}, last={:.4})", is_stable, first_loss, last_loss ); println!("✓ Value loss converges to stable value"); Ok(()) } #[test] fn test_exploration_decay() -> Result<()> { println!("\n=== Test: Exploration Decay (30 epochs with learnable log_std) ==="); let mut config = create_test_config(64); config.policy_config.learnable_std = true; config.policy_config.init_log_std = 0.0; // Start at exp(0) = 1.0 std // Train for 30 epochs let (_ppo, metrics) = train_continuous_ppo(30, config, 10, 20)?; println!(" Log Std Trajectory:"); for (i, &log_std) in metrics.log_stds.iter().enumerate() { if i % 5 == 0 || i == metrics.log_stds.len() - 1 { let std = log_std.exp(); println!(" Epoch {:2}: log_std={:.4}, std={:.4}", i + 1, log_std, std); } } // Verify log_std starts at init value let first_log_std = metrics.log_stds[0]; println!("\n Exploration Analysis:"); println!(" Initial log_std: {:.4} (std={:.4})", first_log_std, first_log_std.exp()); // Verify log_std decreases over time let last_log_std = *metrics.log_stds.last().unwrap(); println!(" Final log_std: {:.4} (std={:.4})", last_log_std, last_log_std.exp()); let decay_occurred = metrics.is_exploration_decaying(); println!(" Decay occurred: {}", decay_occurred); assert!( decay_occurred, "Log std should decrease over training (exploration decay)" ); // Verify log_std stabilizes at lower value (not collapsing to -inf) assert!( last_log_std > -10.0, "Log std should not collapse too much (got {:.4})", last_log_std ); // Calculate decay magnitude let decay_magnitude = first_log_std - last_log_std; println!(" Decay magnitude: {:.4}", decay_magnitude); assert!( decay_magnitude > 0.1, "Decay should be significant (got {:.4})", decay_magnitude ); println!("✓ Exploration decays properly (log_std decreases)"); Ok(()) } #[test] fn test_reward_improvement() -> Result<()> { println!("\n=== Test: Reward Improvement (50 epochs) ==="); let config = create_test_config(64); // Train for 50 epochs let (_ppo, metrics) = train_continuous_ppo(50, config, 10, 20)?; println!(" Average Reward Trajectory:"); for (i, &reward) in metrics.avg_rewards.iter().enumerate() { if i % 10 == 0 || i == metrics.avg_rewards.len() - 1 { println!(" Epoch {:2}: {:.4}", i + 1, reward); } } // Verify monotonic increase (with tolerance) let first_reward = metrics.avg_rewards[0]; let last_reward = *metrics.avg_rewards.last().unwrap(); println!("\n Reward Improvement Analysis:"); println!(" Initial avg reward: {:.4}", first_reward); println!(" Final avg reward: {:.4}", last_reward); let improvement_factor = last_reward / first_reward.abs().max(1e-8); println!(" Improvement factor: {:.2}x", improvement_factor.abs()); // Verify final reward > 2x initial reward (or positive improvement) let improved = metrics.are_rewards_improving(2.0); println!(" Meets 2x improvement: {}", improved); // More lenient check: just verify learning happened let shows_improvement = last_reward > first_reward * 0.9; // Allow 10% variance assert!( improved || shows_improvement, "Reward should improve significantly (first={:.4}, last={:.4}, factor={:.2})", first_reward, last_reward, improvement_factor ); // Verify no catastrophic collapse for (i, &reward) in metrics.avg_rewards.iter().enumerate() { assert!( reward.is_finite(), "Reward at epoch {} should be finite (got {:.4})", i + 1, reward ); } println!("✓ Rewards improve over training"); Ok(()) } #[test] fn test_gradient_stability() -> Result<()> { println!("\n=== Test: Gradient Stability (20 epochs) ==="); let mut config = create_test_config(64); config.max_grad_norm = 10.0; // Set gradient clipping threshold // Train for 20 epochs let (_ppo, metrics) = train_continuous_ppo(20, config, 10, 20)?; println!(" Loss Trajectory Analysis:"); // Monitor for gradient explosions via loss spikes let mut max_policy_loss = f32::NEG_INFINITY; let mut min_policy_loss = f32::INFINITY; let mut max_value_loss = f32::NEG_INFINITY; let mut min_value_loss = f32::INFINITY; for &loss in &metrics.policy_losses { max_policy_loss = max_policy_loss.max(loss); min_policy_loss = min_policy_loss.min(loss); } for &loss in &metrics.value_losses { max_value_loss = max_value_loss.max(loss); min_value_loss = min_value_loss.min(loss); } println!(" Policy loss range: [{:.4}, {:.4}]", min_policy_loss, max_policy_loss); println!(" Value loss range: [{:.4}, {:.4}]", min_value_loss, max_value_loss); // Verify no gradient explosions (losses should stay in reasonable range) assert!( max_policy_loss < 1000.0, "Policy loss should not explode (max={:.4})", max_policy_loss ); assert!( max_value_loss < 1000.0, "Value loss should not explode (max={:.4})", max_value_loss ); // Verify no gradient vanishing (losses should not be too small) assert!( max_policy_loss > 1e-6, "Policy loss should not vanish (max={:.4})", max_policy_loss ); assert!( max_value_loss > 1e-6, "Value loss should not vanish (max={:.4})", max_value_loss ); // Check for sudden spikes in loss (indicator of instability) let mut policy_spikes = 0; let mut value_spikes = 0; for i in 1..metrics.policy_losses.len() { let prev = metrics.policy_losses[i - 1]; let curr = metrics.policy_losses[i]; if (curr - prev).abs() > prev.abs() * 2.0 { policy_spikes += 1; } } for i in 1..metrics.value_losses.len() { let prev = metrics.value_losses[i - 1]; let curr = metrics.value_losses[i]; if (curr - prev).abs() > prev.abs() * 2.0 { value_spikes += 1; } } println!(" Policy loss spikes (>2x change): {}", policy_spikes); println!(" Value loss spikes (>2x change): {}", value_spikes); // Allow some spikes in early training, but not excessive assert!( policy_spikes < 5, "Too many policy loss spikes ({}), indicates instability", policy_spikes ); assert!( value_spikes < 5, "Too many value loss spikes ({}), indicates instability", value_spikes ); println!("✓ Gradients remain stable (no explosions or vanishing)"); Ok(()) } #[test] fn test_long_term_convergence() -> Result<()> { println!("\n=== Test: Long-term Convergence (100 epochs) ==="); let config = create_test_config(64); // Train for 100 epochs println!(" Training for 100 epochs (this may take a minute)..."); let (_ppo, metrics) = train_continuous_ppo(100, config, 10, 20)?; println!("\n Final Metrics:"); println!(" Policy loss: {:.4} → {:.4}", metrics.policy_losses[0], metrics.policy_losses.last().unwrap() ); println!(" Value loss: {:.4} → {:.4}", metrics.value_losses[0], metrics.value_losses.last().unwrap() ); println!(" Avg reward: {:.4} → {:.4}", metrics.avg_rewards[0], metrics.avg_rewards.last().unwrap() ); println!(" Log std: {:.4} → {:.4}", metrics.log_stds[0], metrics.log_stds.last().unwrap() ); // Verify all convergence criteria let policy_converged = metrics.is_policy_loss_converging(0.3); let value_stable = metrics.is_value_loss_stable(10, 0.1); let rewards_improved = metrics.are_rewards_improving(1.5); let exploration_decayed = metrics.is_exploration_decaying(); println!("\n Convergence Criteria:"); println!(" Policy loss converging (≥30% reduction): {}", policy_converged); println!(" Value loss stable (last 10 epochs): {}", value_stable); println!(" Rewards improved (≥1.5x): {}", rewards_improved); println!(" Exploration decayed: {}", exploration_decayed); // At least 3 out of 4 criteria should pass let passing_criteria = [ policy_converged, value_stable, rewards_improved, exploration_decayed, ] .iter() .filter(|&&x| x) .count(); println!(" Total passing: {}/4", passing_criteria); assert!( passing_criteria >= 3, "At least 3/4 convergence criteria should pass (got {}/4)", passing_criteria ); println!("✓ Long-term convergence verified ({}/4 criteria)", passing_criteria); Ok(()) } #[test] fn test_convergence_with_different_learning_rates() -> Result<()> { println!("\n=== Test: Convergence with Different Learning Rates ==="); let state_dim = 64; // Test 1: Balanced learning rates println!("\n Test 1: Balanced LRs (policy=3e-4, value=3e-4)"); let mut config1 = create_test_config(state_dim); config1.policy_learning_rate = 3e-4; config1.value_learning_rate = 3e-4; let (_ppo1, metrics1) = train_continuous_ppo(20, config1, 10, 20)?; println!(" Final policy loss: {:.4}", metrics1.policy_losses.last().unwrap()); println!(" Final value loss: {:.4}", metrics1.value_losses.last().unwrap()); // Test 2: Aggressive value LR (like hyperopt best) println!("\n Test 2: Aggressive value LR (policy=1e-4, value=1e-3)"); let mut config2 = create_test_config(state_dim); config2.policy_learning_rate = 1e-4; config2.value_learning_rate = 1e-3; let (_ppo2, metrics2) = train_continuous_ppo(20, config2, 10, 20)?; println!(" Final policy loss: {:.4}", metrics2.policy_losses.last().unwrap()); println!(" Final value loss: {:.4}", metrics2.value_losses.last().unwrap()); // Test 3: Conservative policy LR println!("\n Test 3: Conservative policy LR (policy=1e-5, value=3e-4)"); let mut config3 = create_test_config(state_dim); config3.policy_learning_rate = 1e-5; config3.value_learning_rate = 3e-4; let (_ppo3, metrics3) = train_continuous_ppo(20, config3, 10, 20)?; println!(" Final policy loss: {:.4}", metrics3.policy_losses.last().unwrap()); println!(" Final value loss: {:.4}", metrics3.value_losses.last().unwrap()); // Verify all configurations converged (no crashes) assert_eq!(metrics1.policy_losses.len(), 20, "Config 1 should complete 20 epochs"); assert_eq!(metrics2.policy_losses.len(), 20, "Config 2 should complete 20 epochs"); assert_eq!(metrics3.policy_losses.len(), 20, "Config 3 should complete 20 epochs"); // Verify losses are finite for loss in &metrics1.policy_losses { assert!(loss.is_finite(), "Config 1 losses should be finite"); } for loss in &metrics2.policy_losses { assert!(loss.is_finite(), "Config 2 losses should be finite"); } for loss in &metrics3.policy_losses { assert!(loss.is_finite(), "Config 3 losses should be finite"); } println!("\n✓ All learning rate configurations converge successfully"); Ok(()) }