//! Unit tests for PPO Continuous Policy //! //! Tests continuous action spaces, policy network, and action sampling. use anyhow::Result; use candle_core::{Device, Tensor}; use ml::ppo::continuous_policy::{ContinuousPolicyConfig, ContinuousPolicyNetwork}; #[test] fn test_continuous_policy_creation() -> Result<()> { let device = Device::Cpu; let config = ContinuousPolicyConfig { input_dim: 64, hidden_dim: 128, action_dim: 4, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; // Verify policy was created assert!(policy.input_dim() == 64); assert!(policy.action_dim() == 4); Ok(()) } #[test] fn test_continuous_policy_forward() -> Result<()> { let device = Device::Cpu; let batch_size = 8; let input_dim = 64; let action_dim = 4; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 128, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; // Create state input let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; // Forward pass returns (mean, std) let (mean, std) = policy.forward(&state)?; // Verify shapes assert_eq!(mean.dims(), &[batch_size, action_dim]); assert_eq!(std.dims(), &[batch_size, action_dim]); // Verify std is positive let std_min = std.min(0)?.min(0)?.to_scalar::()?; assert!(std_min > 0.0, "Standard deviation should be positive"); Ok(()) } #[test] fn test_continuous_policy_action_sampling() -> Result<()> { let device = Device::Cpu; let batch_size = 4; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; // Sample actions let action = policy.sample_action(&state)?; // Verify action shape assert_eq!(action.dims(), &[batch_size, action_dim]); // Actions should be finite let action_max = action.abs()?.max(0)?.max(0)?.to_scalar::()?; assert!(action_max.is_finite()); Ok(()) } #[test] fn test_continuous_policy_deterministic_mode() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::ones((batch_size, input_dim), candle_core::DType::F32, &device)?; // In deterministic mode, should return mean let (mean, _) = policy.forward(&state)?; let action = policy.deterministic_action(&state)?; // Action should equal mean in deterministic mode let diff = (&action - &mean)?.abs()?.sum_all()?.to_scalar::()?; assert!(diff < 1e-5, "Deterministic action should equal mean"); Ok(()) } #[test] fn test_continuous_policy_log_prob() -> Result<()> { let device = Device::Cpu; let batch_size = 4; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; let action = Tensor::randn(0.0f32, 1.0, (batch_size, action_dim), &device)?; // Compute log probability let log_prob = policy.log_prob(&state, &action)?; // Verify shape assert_eq!(log_prob.dims(), &[batch_size]); // Log probabilities should be negative or zero let log_prob_max = log_prob.max(0)?.to_scalar::()?; assert!(log_prob_max <= 0.01, "Log probabilities should be ≤ 0"); Ok(()) } #[test] fn test_continuous_policy_entropy() -> Result<()> { let device = Device::Cpu; let batch_size = 4; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; // Compute entropy let entropy = policy.entropy(&state)?; // Entropy should be positive (Gaussian entropy > 0) let entropy_min = entropy.min(0)?.to_scalar::()?; assert!(entropy_min > 0.0, "Entropy should be positive"); Ok(()) } #[test] fn test_continuous_policy_gradient_flow() -> Result<()> { let device = Device::Cpu; let batch_size = 4; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; let action = policy.sample_action(&state)?; // Compute log prob and loss let log_prob = policy.log_prob(&state, &action)?; let loss = log_prob.sum_all()?; // Verify backward pass works loss.backward()?; Ok(()) } #[test] fn test_continuous_policy_action_bounds() -> Result<()> { let device = Device::Cpu; let batch_size = 10; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; // Sample many actions for _ in 0..100 { let action = policy.sample_action(&state)?; // Actions should be within reasonable bounds (e.g., ±10) let action_max = action.abs()?.max(0)?.max(0)?.to_scalar::()?; assert!(action_max < 100.0, "Actions should not be extreme"); } Ok(()) } #[test] fn test_continuous_policy_different_action_dims() -> Result<()> { let device = Device::Cpu; let batch_size = 4; let input_dim = 64; // Test various action dimensions for action_dim in [1, 2, 4, 8] { let config = ContinuousPolicyConfig { input_dim, hidden_dim: 128, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::randn(0.0f32, 1.0, (batch_size, input_dim), &device)?; let (mean, std) = policy.forward(&state)?; assert_eq!(mean.dim(1)?, action_dim); assert_eq!(std.dim(1)?, action_dim); } Ok(()) } #[test] fn test_continuous_policy_consistency() -> Result<()> { let device = Device::Cpu; let batch_size = 2; let input_dim = 32; let action_dim = 2; let config = ContinuousPolicyConfig { input_dim, hidden_dim: 64, action_dim, learning_rate: 3e-4, }; let policy = ContinuousPolicyNetwork::new(config, &device)?; let state = Tensor::ones((batch_size, input_dim), candle_core::DType::F32, &device)?; // Same input should give same mean (deterministic forward) let (mean1, _) = policy.forward(&state)?; let (mean2, _) = policy.forward(&state)?; let diff = (&mean1 - &mean2)?.abs()?.sum_all()?.to_scalar::()?; assert!(diff < 1e-6, "Forward pass should be deterministic"); Ok(()) }