- 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)
265 lines
8.9 KiB
Rust
265 lines
8.9 KiB
Rust
/// 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(())
|
||
}
|