Files
foxhunt/ml/tests/dqn_e2e_training_validation_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

265 lines
8.9 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
/// DQN E2E Training Validation Test
///
/// This test validates that the DQN network ACTUALLY LEARNS during training by:
/// 1. Training for 10 epochs
/// 2. Extracting Q-values at epoch 1 and epoch 10
/// 3. Asserting Q-values CHANGED (not constant)
/// 4. Asserting gradients are non-zero and different across epochs
///
/// This test will FAIL if the network isn't learning (constant Q-values bug).
use ml::dqn::dqn::WorkingDQN;
use ml::dqn::dqn::WorkingDQNConfig;
use ml::MLError;
#[test]
fn test_dqn_actually_learns() -> Result<(), MLError> {
// Initialize DQN agent with minimal config
let mut config = WorkingDQNConfig::conservative();
config.num_actions = 45; // 45-action space (5×3×3)
config.state_dim = 54; // Standard state dimension
config.learning_rate = 1e-4;
config.batch_size = 32;
config.replay_buffer_capacity = 1000;
config.warmup_steps = 100; // Minimal warmup for test
let mut agent = WorkingDQN::new(config)?;
// Generate dummy state
let dummy_state = vec![0.1f32; 54];
// ===== EPOCH 1: Initial Q-values =====
// Train for 1 epoch (collect experiences during warmup)
for _ in 0..100 {
let action = agent.select_action(&dummy_state)?;
let next_state = dummy_state.clone();
// Create experience with small reward
let experience = ml::dqn::experience::Experience::new(
dummy_state.clone(),
action.to_index() as u8,
0.01, // Small positive reward
next_state,
false, // Not terminal
);
agent.store_experience(experience)?;
}
// Now we can train (warmup complete)
let mut epoch1_q_values = Vec::new();
let device = candle_core::Device::cuda_if_available(0).unwrap_or(candle_core::Device::Cpu);
for _ in 0..10 {
if let Ok((loss, grad_norm)) = agent.train_step(None) {
// Extract Q-values for all actions
let state_tensor = candle_core::Tensor::from_vec(
dummy_state.clone(),
(1, 54),
&device,
).map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))?;
let q_values = agent.forward(&state_tensor)?;
let q_vec: Vec<f32> = q_values.to_vec2::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to extract Q-values: {}", e)))?
.into_iter()
.flatten()
.collect();
epoch1_q_values.push(q_vec);
println!("Epoch 1, step {}: loss={:.6}, grad_norm={:.6}", epoch1_q_values.len(), loss, grad_norm);
}
}
// ===== EPOCHS 2-10: Continue training =====
for epoch in 2..=10 {
// Collect more experiences
for _ in 0..20 {
let action = agent.select_action(&dummy_state)?;
let next_state = dummy_state.clone();
let experience = ml::dqn::experience::Experience::new(
dummy_state.clone(),
action.to_index() as u8,
0.01,
next_state,
false,
);
agent.store_experience(experience)?;
}
// Train
for _ in 0..10 {
if let Ok(_) = agent.train_step(None) {
// Continue training
}
}
println!("Completed epoch {}", epoch);
}
// ===== EPOCH 10: Final Q-values =====
let mut epoch10_q_values = Vec::new();
let device = candle_core::Device::cuda_if_available(0).unwrap_or(candle_core::Device::Cpu);
for _ in 0..10 {
if let Ok(_) = agent.train_step(None) {
let state_tensor = candle_core::Tensor::from_vec(
dummy_state.clone(),
(1, 54),
&device,
).map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))?;
let q_values = agent.forward(&state_tensor)?;
let q_vec: Vec<f32> = q_values.to_vec2::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to extract Q-values: {}", e)))?
.into_iter()
.flatten()
.collect();
epoch10_q_values.push(q_vec);
}
}
// ===== ASSERTIONS: Q-values MUST CHANGE =====
assert!(!epoch1_q_values.is_empty(), "No Q-values collected at epoch 1");
assert!(!epoch10_q_values.is_empty(), "No Q-values collected at epoch 10");
// Compare first and last Q-values
let epoch1_first = &epoch1_q_values[0];
let epoch10_last = &epoch10_q_values[epoch10_q_values.len() - 1];
assert_eq!(epoch1_first.len(), 45, "Q-values should have 45 dimensions (not 3!)");
assert_eq!(epoch10_last.len(), 45, "Q-values should have 45 dimensions (not 3!)");
// Calculate change in Q-values
let mut total_change = 0.0f32;
let mut max_change = 0.0f32;
for i in 0..45 {
let change = (epoch10_last[i] - epoch1_first[i]).abs();
total_change += change;
max_change = max_change.max(change);
}
let avg_change = total_change / 45.0;
println!("\n===== Q-VALUE CHANGE ANALYSIS =====");
println!("Average Q-value change: {:.6}", avg_change);
println!("Maximum Q-value change: {:.6}", max_change);
println!("Epoch 1 Q-value range: [{:.6}, {:.6}]",
epoch1_first.iter().cloned().fold(f32::INFINITY, f32::min),
epoch1_first.iter().cloned().fold(f32::NEG_INFINITY, f32::max));
println!("Epoch 10 Q-value range: [{:.6}, {:.6}]",
epoch10_last.iter().cloned().fold(f32::INFINITY, f32::min),
epoch10_last.iter().cloned().fold(f32::NEG_INFINITY, f32::max));
// CRITICAL ASSERTIONS
assert!(
avg_change > 0.01,
"FAIL: Q-values barely changed (avg change: {:.6}). Network is not learning!",
avg_change
);
assert!(
max_change > 0.1,
"FAIL: Maximum Q-value change too small (max change: {:.6}). Network is not learning!",
max_change
);
// Check that not all Q-values are the same (constant bug)
let epoch10_std_dev = {
let mean = epoch10_last.iter().sum::<f32>() / 45.0;
let variance = epoch10_last.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / 45.0;
variance.sqrt()
};
assert!(
epoch10_std_dev > 0.01,
"FAIL: Q-values are constant (std dev: {:.6}). All Q-values are the same!",
epoch10_std_dev
);
println!("✓ TEST PASSED: Network is learning (Q-values changed significantly)");
Ok(())
}
#[test]
fn test_q_values_not_hardcoded() -> Result<(), MLError> {
// Test that Q-values are NOT hardcoded to [80, 750, -3200] pattern
let mut config = WorkingDQNConfig::conservative();
config.num_actions = 45;
config.state_dim = 54;
let agent = WorkingDQN::new(config)?;
// Generate two different states
let state1 = vec![0.1f32; 54];
let state2 = vec![0.9f32; 54];
// Get Q-values for both states
let device = candle_core::Device::cuda_if_available(0).unwrap_or(candle_core::Device::Cpu);
let state1_tensor = candle_core::Tensor::from_vec(
state1.clone(),
(1, 54),
&device,
).map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))?;
let state2_tensor = candle_core::Tensor::from_vec(
state2.clone(),
(1, 54),
&device,
).map_err(|e| MLError::ModelError(format!("Failed to create state tensor: {}", e)))?;
let q1 = agent.forward(&state1_tensor)?;
let q2 = agent.forward(&state2_tensor)?;
let q1_vec: Vec<f32> = q1.to_vec2::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to extract Q-values: {}", e)))?
.into_iter()
.flatten()
.collect();
let q2_vec: Vec<f32> = q2.to_vec2::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to extract Q-values: {}", e)))?
.into_iter()
.flatten()
.collect();
// Check that Q-values are different for different states
let mut diff_count = 0;
for i in 0..45 {
if (q1_vec[i] - q2_vec[i]).abs() > 0.0001 {
diff_count += 1;
}
}
assert!(
diff_count > 20, // At least half should be different
"FAIL: Q-values are too similar for different states ({}/45 different). Network may be outputting hardcoded values!",
diff_count
);
// Check NOT matching the [80, 750, -3200] pattern
// (This pattern appears when logging only indices 0, 1, 2 of 45-action network)
let suspicious_pattern = (q1_vec[0] > 50.0 && q1_vec[0] < 120.0) &&
(q1_vec[1] > 700.0 && q1_vec[1] < 800.0) &&
(q1_vec[2] < -3000.0 && q1_vec[2] > -3500.0);
assert!(
!suspicious_pattern,
"FAIL: Q-values match suspicious hardcoded pattern [~80, ~750, ~-3200]! First 3 Q-values: [{:.2}, {:.2}, {:.2}]",
q1_vec[0], q1_vec[1], q1_vec[2]
);
println!("✓ TEST PASSED: Q-values are network-computed (not hardcoded)");
Ok(())
}