Files
foxhunt/ml/tests/regime_conditional_dqn_test.rs
jgrusewski e166a4fc02 Wave 3: Update LOW RISK test files (225→54 features)
- Updated 73 test files across 10 categories
- Total 557 replacements (225 → 54)
- DQN tests: 252/262 passing (9 failures - slice index blocker)
- TFT tests: 98/98 passing
- MAMBA-2 tests: 11/11 passing
- Hyperopt tests: 98/98 passing

Critical findings:
- Blocker: ml/src/trainers/dqn.rs:3444 hardcoded slice indices
- Architecture mismatch: extract_current_features() vs extract_current_features_v2()

Wave 3 Agent breakdown:
- Agent 1: DQN test files (12 files)
- Agent 2: PPO test files (2 files)
- Agent 3: TFT test files (6 files)
- Agent 4: MAMBA-2 test files (2 files)
- Agent 5: Feature extraction tests (3 files)
- Agent 6: Integration test files (9 files)
- Agent 7: Data loader test files (3 files)
- Agent 8: Hyperopt test files (1 file)
- Agent 9: Benchmark test files (9 files)
- Agent 10: Utility & misc test files (73 files)

Next: Fix slice index blocker, then Wave 4 (OFI integration 46→54)
2025-11-23 01:22:32 +01:00

478 lines
15 KiB
Rust

//! 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; 54];
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: 54,
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, 54], 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::<f32>()?;
let q_ranging_vals = q_ranging.flatten_all()?.to_vec1::<f32>()?;
// 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 = 54;
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; 54];
let mut next_state = vec![0.0_f32; 54];
// 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 = 54;
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; 54];
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, 54], candle_core::DType::F32, &device)?;
let q_before_trending = dqn.forward(&state, RegimeType::Trending)?.flatten_all()?.to_vec1::<f32>()?;
let q_before_ranging = dqn.forward(&state, RegimeType::Ranging)?.flatten_all()?.to_vec1::<f32>()?;
let q_before_volatile = dqn.forward(&state, RegimeType::Volatile)?.flatten_all()?.to_vec1::<f32>()?;
// 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::<f32>()?;
let q_after_ranging = dqn.forward(&state, RegimeType::Ranging)?.flatten_all()?.to_vec1::<f32>()?;
let q_after_volatile = dqn.forward(&state, RegimeType::Volatile)?.flatten_all()?.to_vec1::<f32>()?;
// 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::<f32>()?;
// 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::<f32>()?;
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; 54];
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 = 54;
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; 54];
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, 54], 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; 54];
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; 54];
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(())
}