CRITICAL FINDINGS from 3-trial validation: - 85,120 gradient clipping warnings (81.6% of logs) - REGRESSION - Rainbow features DISABLED: use_dueling=false, use_distributional=false, use_noisy_nets=false - Negative Q-values confirmed: HOLD -1000 to -3250 - Performance: Sharpe 0.29 (target 0.77) Changes: - Fixed N-Step compilation (7/7 tests passing) - Fixed Distributional compilation (6/6 tests passing) - Fixed Dueling CUDA errors (10/10 tests passing) - Added TDD validation for state_dim=225 - Total: 23/23 Wave 11 tests passing (100%) Issues requiring investigation: 1. Why are Dueling/Distributional/Noisy disabled in hyperopt? 2. Why gradient explosion despite previous fixes? 3. Test coverage gaps - unit tests pass but integration fails 🤖 Generated with Claude Code Co-Authored-By: Claude <noreply@anthropic.com>
453 lines
15 KiB
Rust
453 lines
15 KiB
Rust
//! 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<f32> = metrics
|
|
.value_losses
|
|
.iter()
|
|
.rev()
|
|
.take(window)
|
|
.copied()
|
|
.collect();
|
|
let mean = recent.iter().sum::<f32>() / recent.len() as f32;
|
|
let variance = recent.iter().map(|&x| (x - mean).powi(2)).sum::<f32>() / 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(())
|
|
}
|