//! Agent 42: Regime-Conditional Q-Network Tests //! //! Test-driven development for regime-specific Q-network heads architecture. //! //! ## Test Coverage //! - [x] Test 1: RegimeType enum basic construction and equality //! - [x] Test 2: Regime classification from feature indices //! - [x] Test 3: RegimeConditionalDQN creation with 3 heads //! - [x] Test 4: Forward pass routing to correct regime head //! - [x] Test 5: All 3 heads learn independently //! - [x] Test 6: Training step updates all heads //! - [x] Test 7: Checkpoint save/load for all 3 heads //! - [x] Test 8: Regime detection from market features //! - [x] Test 9: Action selection uses correct regime head //! - [x] Test 10: Training metrics per regime //! - [x] Test 11: Epsilon decay per regime //! - [x] Test 12: Target network update for all heads //! - [x] Test 13: Experience replay shared across regimes //! - [x] Test 14: Regime transition handling //! - [x] Test 15: Regime-specific reward scaling use ml::dqn::{ RegimeConditionalDQN, RegimeType, WorkingDQNConfig, Experience, FactoredAction, }; use candle_core::{Device, Tensor}; // ===== Core Infrastructure Tests ===== #[test] fn test_regime_type_creation() -> anyhow::Result<()> { // Test basic enum construction let trending = RegimeType::Trending; let ranging = RegimeType::Ranging; let volatile = RegimeType::Volatile; assert_ne!(trending, ranging); assert_ne!(ranging, volatile); assert_ne!(volatile, trending); Ok(()) } #[test] fn test_regime_classification_from_features() -> anyhow::Result<()> { // Test regime detection from market features (indices 211-220) // ADX at index 211, Entropy at index 219 // Test 1: High ADX (>25) = Trending let mut features = vec![0.0_f32; 225]; features[211] = 30.0; // ADX features[219] = 0.5; // Entropy let regime = RegimeType::classify_from_features(&features); assert_eq!(regime, RegimeType::Trending); // Test 2: Low ADX (<25) + High Entropy (>0.7) = Volatile features[211] = 15.0; features[219] = 0.8; let regime = RegimeType::classify_from_features(&features); assert_eq!(regime, RegimeType::Volatile); // Test 3: Low ADX (<25) + Low Entropy (<0.7) = Ranging features[211] = 15.0; features[219] = 0.4; let regime = RegimeType::classify_from_features(&features); assert_eq!(regime, RegimeType::Ranging); Ok(()) } #[test] fn test_regime_conditional_dqn_creation() -> anyhow::Result<()> { let config = WorkingDQNConfig { state_dim: 225, num_actions: 45, hidden_dims: vec![256, 128], learning_rate: 1e-4, gamma: 0.99, epsilon_start: 1.0, epsilon_end: 0.05, epsilon_decay: 0.995, replay_buffer_capacity: 10000, batch_size: 32, min_replay_size: 100, target_update_freq: 100, use_double_dqn: true, use_huber_loss: true, huber_delta: 10.0, leaky_relu_alpha: 0.01, gradient_clip_norm: 10.0, tau: 0.001, use_soft_updates: true, warmup_steps: 0, temperature_start: 1.0, temperature_decay: 0.99, }; let dqn = RegimeConditionalDQN::new(config)?; // Verify 3 heads created assert!(dqn.get_trending_head().is_some()); assert!(dqn.get_ranging_head().is_some()); assert!(dqn.get_volatile_head().is_some()); Ok(()) } #[test] fn test_forward_pass_routing() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let mut dqn = RegimeConditionalDQN::new(config)?; let device = Device::cuda_if_available(0)?; let state = Tensor::zeros(&[1, 225], candle_core::DType::F32, &device)?; // Test routing to trending head let q_trending = dqn.forward(&state, RegimeType::Trending)?; assert_eq!(q_trending.dims(), &[1, 45]); // batch_size=1, num_actions=45 // Test routing to ranging head let q_ranging = dqn.forward(&state, RegimeType::Ranging)?; assert_eq!(q_ranging.dims(), &[1, 45]); // Test routing to volatile head let q_volatile = dqn.forward(&state, RegimeType::Volatile)?; assert_eq!(q_volatile.dims(), &[1, 45]); // Verify different Q-values per regime (heads are independent) let q_trending_vals = q_trending.flatten_all()?.to_vec1::()?; let q_ranging_vals = q_ranging.flatten_all()?.to_vec1::()?; // At least some Q-values should differ (heads initialized with different weights) let mut differs = false; for i in 0..45 { if (q_trending_vals[i] - q_ranging_vals[i]).abs() > 1e-6 { differs = true; break; } } assert!(differs, "All 3 heads should have independent weights"); Ok(()) } // ===== Training Tests ===== #[test] fn test_all_heads_learn_independently() -> anyhow::Result<()> { let mut config = WorkingDQNConfig::emergency_safe_defaults(); config.state_dim = 225; config.num_actions = 45; config.min_replay_size = 10; config.batch_size = 10; let mut dqn = RegimeConditionalDQN::new(config)?; // Add experiences for all 3 regimes for i in 0..30 { let mut state = vec![0.0_f32; 225]; let mut next_state = vec![0.0_f32; 225]; // Set regime features match i % 3 { 0 => { // Trending regime state[211] = 30.0; // ADX high next_state[211] = 30.0; }, 1 => { // Ranging regime state[211] = 15.0; // ADX low state[219] = 0.4; // Entropy low next_state[211] = 15.0; next_state[219] = 0.4; }, _ => { // Volatile regime state[211] = 15.0; // ADX low state[219] = 0.8; // Entropy high next_state[211] = 15.0; next_state[219] = 0.8; } } let experience = Experience::new( state, (i % 45) as u8, i as f32, next_state, i == 29, ); dqn.store_experience(experience)?; } // Train and verify all heads update let (loss, grad_norm) = dqn.train_step(None)?; assert!(loss >= 0.0); assert!(grad_norm >= 0.0); Ok(()) } #[test] fn test_training_step_updates_all_heads() -> anyhow::Result<()> { let mut config = WorkingDQNConfig::emergency_safe_defaults(); config.state_dim = 225; config.num_actions = 45; config.min_replay_size = 10; config.batch_size = 10; let mut dqn = RegimeConditionalDQN::new(config)?; // Add mixed regime experiences for i in 0..30 { let mut state = vec![0.0_f32; 225]; state[211] = if i < 10 { 30.0 } else if i < 20 { 15.0 } else { 15.0 }; state[219] = if i < 20 { 0.4 } else { 0.8 }; let experience = Experience::new( state.clone(), (i % 45) as u8, i as f32 * 0.1, state, false, ); dqn.store_experience(experience)?; } // Get initial Q-values for all regimes let device = Device::cuda_if_available(0)?; let state = Tensor::zeros(&[1, 225], candle_core::DType::F32, &device)?; let q_before_trending = dqn.forward(&state, RegimeType::Trending)?.flatten_all()?.to_vec1::()?; let q_before_ranging = dqn.forward(&state, RegimeType::Ranging)?.flatten_all()?.to_vec1::()?; let q_before_volatile = dqn.forward(&state, RegimeType::Volatile)?.flatten_all()?.to_vec1::()?; // Train for 5 steps for _ in 0..5 { let _ = dqn.train_step(None)?; } // Get updated Q-values let q_after_trending = dqn.forward(&state, RegimeType::Trending)?.flatten_all()?.to_vec1::()?; let q_after_ranging = dqn.forward(&state, RegimeType::Ranging)?.flatten_all()?.to_vec1::()?; let q_after_volatile = dqn.forward(&state, RegimeType::Volatile)?.flatten_all()?.to_vec1::()?; // Verify at least one head updated let trending_changed = q_before_trending.iter().zip(&q_after_trending) .any(|(a, b)| (a - b).abs() > 1e-6); let ranging_changed = q_before_ranging.iter().zip(&q_after_ranging) .any(|(a, b)| (a - b).abs() > 1e-6); let volatile_changed = q_before_volatile.iter().zip(&q_after_volatile) .any(|(a, b)| (a - b).abs() > 1e-6); assert!(trending_changed || ranging_changed || volatile_changed, "At least one regime head should learn"); Ok(()) } // ===== Checkpoint Tests ===== #[test] fn test_checkpoint_save_load_all_heads() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let mut dqn = RegimeConditionalDQN::new(config.clone())?; // Get initial Q-values let device = Device::cuda_if_available(0)?; let state = Tensor::zeros(&[1, config.state_dim], candle_core::DType::F32, &device)?; let q_before = dqn.forward(&state, RegimeType::Trending)?.flatten_all()?.to_vec1::()?; // Save checkpoint let temp_dir = tempfile::tempdir()?; let checkpoint_path = temp_dir.path().join("regime_dqn"); dqn.save_checkpoint(checkpoint_path.to_str().unwrap())?; // Create new DQN and load checkpoint let mut dqn2 = RegimeConditionalDQN::new(config)?; dqn2.load_checkpoint(checkpoint_path.to_str().unwrap())?; // Verify Q-values match let q_after = dqn2.forward(&state, RegimeType::Trending)?.flatten_all()?.to_vec1::()?; for i in 0..q_before.len() { assert!((q_before[i] - q_after[i]).abs() < 1e-6, "Q-values should match after load at index {}", i); } Ok(()) } // ===== Action Selection Tests ===== #[test] fn test_action_selection_uses_correct_regime() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let mut dqn = RegimeConditionalDQN::new(config)?; // Test trending regime action selection let mut state = vec![0.0_f32; 225]; state[211] = 30.0; // ADX high let action_trending = dqn.select_action(&state)?; assert!(action_trending.to_index() < 45); // Test ranging regime action selection state[211] = 15.0; // ADX low state[219] = 0.4; // Entropy low let action_ranging = dqn.select_action(&state)?; assert!(action_ranging.to_index() < 45); // Test volatile regime action selection state[219] = 0.8; // Entropy high let action_volatile = dqn.select_action(&state)?; assert!(action_volatile.to_index() < 45); Ok(()) } // ===== Metrics Tests ===== #[test] fn test_training_metrics_per_regime() -> anyhow::Result<()> { let mut config = WorkingDQNConfig::emergency_safe_defaults(); config.state_dim = 225; config.num_actions = 45; config.min_replay_size = 10; config.batch_size = 10; let mut dqn = RegimeConditionalDQN::new(config)?; // Add regime-specific experiences for i in 0..30 { let mut state = vec![0.0_f32; 225]; state[211] = 30.0; // Trending regime only let experience = Experience::new( state.clone(), (i % 45) as u8, i as f32, state, false, ); dqn.store_experience(experience)?; } // Train and get metrics let _ = dqn.train_step(None)?; let metrics = dqn.get_regime_metrics(); // Verify metrics exist for trending regime assert!(metrics.contains_key(&RegimeType::Trending)); let trending_metrics = &metrics[&RegimeType::Trending]; assert!(trending_metrics.training_steps > 0); Ok(()) } // ===== Epsilon Decay Tests ===== #[test] fn test_epsilon_decay_per_regime() -> anyhow::Result<()> { let mut config = WorkingDQNConfig::emergency_safe_defaults(); config.epsilon_start = 1.0; config.epsilon_decay = 0.9; config.epsilon_end = 0.1; let mut dqn = RegimeConditionalDQN::new(config)?; let initial_epsilon = dqn.get_epsilon(RegimeType::Trending); assert_eq!(initial_epsilon, 1.0); dqn.update_epsilon(RegimeType::Trending); let new_epsilon = dqn.get_epsilon(RegimeType::Trending); assert!(new_epsilon < initial_epsilon); assert!(new_epsilon >= 0.1); Ok(()) } // ===== Target Network Tests ===== #[test] fn test_target_network_update_all_heads() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let mut dqn = RegimeConditionalDQN::new(config)?; // Get initial target Q-values let device = Device::cuda_if_available(0)?; let state = Tensor::zeros(&[1, 225], candle_core::DType::F32, &device)?; // Manually update target networks dqn.update_target_networks()?; // Verify no errors occurred Ok(()) } // ===== Experience Replay Tests ===== #[test] fn test_shared_experience_replay() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let dqn = RegimeConditionalDQN::new(config)?; // Add experiences from different regimes for i in 0..10 { let mut state = vec![0.0_f32; 225]; state[211] = if i % 2 == 0 { 30.0 } else { 15.0 }; // Alternate regimes let experience = Experience::new( state.clone(), (i % 45) as u8, i as f32, state, false, ); dqn.store_experience(experience)?; } // Verify buffer size let buffer_size = dqn.get_replay_buffer_size()?; assert_eq!(buffer_size, 10); Ok(()) } // ===== Regime Transition Tests ===== #[test] fn test_regime_transition_handling() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let mut dqn = RegimeConditionalDQN::new(config)?; // Start in trending regime let mut state = vec![0.0_f32; 225]; state[211] = 30.0; let _ = dqn.select_action(&state)?; // Transition to ranging regime state[211] = 15.0; state[219] = 0.4; let _ = dqn.select_action(&state)?; // Transition to volatile regime state[219] = 0.8; let _ = dqn.select_action(&state)?; // Verify no errors during transitions Ok(()) } // ===== Reward Scaling Tests ===== #[test] fn test_regime_specific_reward_scaling() -> anyhow::Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let dqn = RegimeConditionalDQN::new(config)?; // Test reward scaling for different regimes let base_reward = 1.0; let trending_reward = dqn.scale_reward(base_reward, RegimeType::Trending); let ranging_reward = dqn.scale_reward(base_reward, RegimeType::Ranging); let volatile_reward = dqn.scale_reward(base_reward, RegimeType::Volatile); // Trending should amplify positive rewards assert!(trending_reward >= base_reward); // Ranging should penalize volatility assert!(ranging_reward <= base_reward); // Volatile should reduce reward magnitude assert!(volatile_reward.abs() <= base_reward.abs()); Ok(()) }