- 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)
187 lines
5.8 KiB
Rust
187 lines
5.8 KiB
Rust
//! Test for DQN Large Network Architecture (256, 128, 64 neurons)
|
||
//!
|
||
//! Validates that the DQN network uses the larger architecture to prevent
|
||
//! gradient collapse and Q-value collapse issues discovered in Wave 9-A4.
|
||
//!
|
||
//! Test-driven development: This test should FAIL before implementation and
|
||
//! PASS after updating the network architecture.
|
||
|
||
use candle_core::Tensor;
|
||
use ml::dqn::dqn::{WorkingDQN, WorkingDQNConfig};
|
||
|
||
#[test]
|
||
fn test_large_network_architecture() -> Result<(), Box<dyn std::error::Error>> {
|
||
// Create config with large network architecture
|
||
let config = WorkingDQNConfig {
|
||
state_dim: 54,
|
||
num_actions: 3,
|
||
hidden_dims: vec![256, 128, 64], // Target architecture
|
||
learning_rate: 0.0001,
|
||
gamma: 0.99,
|
||
epsilon_start: 1.0,
|
||
epsilon_end: 0.01,
|
||
epsilon_decay: 0.995,
|
||
replay_buffer_capacity: 10000,
|
||
batch_size: 32,
|
||
min_replay_size: 100,
|
||
target_update_freq: 1000,
|
||
use_double_dqn: true,
|
||
use_huber_loss: true,
|
||
huber_delta: 1.0,
|
||
leaky_relu_alpha: 0.01,
|
||
};
|
||
|
||
// Create DQN with large network
|
||
let dqn = WorkingDQN::new(config)?;
|
||
|
||
// Test 1: Forward pass with 54-dim input → 3-dim output
|
||
let state = Tensor::randn(0f32, 1.0, (1, 54), dqn.device())?;
|
||
let q_values = dqn.forward(&state)?;
|
||
|
||
// Verify output shape
|
||
assert_eq!(
|
||
q_values.dims(),
|
||
&[1, 3],
|
||
"Q-values should have shape [1, 3] (batch=1, actions=3)"
|
||
);
|
||
|
||
// Test 2: Batch forward pass (validate memory safety)
|
||
let batch_size = 128; // Production batch size
|
||
let batch_state = Tensor::randn(0f32, 1.0, (batch_size, 54), dqn.device())?;
|
||
let batch_q_values = dqn.forward(&batch_state)?;
|
||
|
||
assert_eq!(
|
||
batch_q_values.dims(),
|
||
&[batch_size, 3],
|
||
"Batch Q-values should have shape [128, 3]"
|
||
);
|
||
|
||
// Test 3: Verify Q-values are finite (no NaN/Inf)
|
||
let q_vals_vec = batch_q_values.flatten_all()?.to_vec1::<f32>()?;
|
||
|
||
for (i, &val) in q_vals_vec.iter().enumerate() {
|
||
assert!(
|
||
val.is_finite(),
|
||
"Q-value at index {} is not finite: {}",
|
||
i,
|
||
val
|
||
);
|
||
}
|
||
|
||
println!("✓ Large network architecture test passed");
|
||
println!(" - Network: [256, 128, 64]");
|
||
println!(" - Input: 54 dims");
|
||
println!(" - Output: 3 actions");
|
||
println!(" - Batch size: 128 (validated)");
|
||
|
||
Ok(())
|
||
}
|
||
|
||
#[test]
|
||
fn test_large_network_parameter_count() -> Result<(), Box<dyn std::error::Error>> {
|
||
// Calculate expected parameter count for [256, 128, 64] network
|
||
// Layer 1: 54 × 256 = 57,600
|
||
// Layer 2: 256 × 128 = 32,768
|
||
// Layer 3: 128 × 64 = 8,192
|
||
// Output: 64 × 3 = 192
|
||
// Total: 98,752 parameters
|
||
|
||
let expected_params = 57_600 + 32_768 + 8_192 + 192;
|
||
assert_eq!(
|
||
expected_params, 98_752,
|
||
"Expected parameter count should be 98,752"
|
||
);
|
||
|
||
// Verify this is ~2.5x larger than current [128, 64, 32]
|
||
let old_params = 28_800 + 8_192 + 2_048 + 96;
|
||
assert_eq!(old_params, 39_136, "Old parameter count should be 39,136");
|
||
|
||
let ratio = expected_params as f64 / old_params as f64;
|
||
assert!(
|
||
(ratio - 2.52).abs() < 0.01,
|
||
"Parameter count should increase by ~2.5x, got {:.2}x",
|
||
ratio
|
||
);
|
||
|
||
println!("✓ Parameter count validation passed");
|
||
println!(" - Old network: 39,136 params");
|
||
println!(" - New network: 98,752 params");
|
||
println!(" - Increase: {:.2}x", ratio);
|
||
|
||
Ok(())
|
||
}
|
||
|
||
#[test]
|
||
fn test_large_network_prevents_gradient_collapse() -> Result<(), Box<dyn std::error::Error>> {
|
||
// This test validates that the larger network has enough capacity
|
||
// to prevent gradient collapse (36,341 → 0.80 issue from Wave 9-A4)
|
||
|
||
let config = WorkingDQNConfig {
|
||
state_dim: 54,
|
||
num_actions: 3,
|
||
hidden_dims: vec![256, 128, 64],
|
||
learning_rate: 0.0001,
|
||
gamma: 0.99,
|
||
epsilon_start: 0.1, // Low epsilon for greedy actions
|
||
epsilon_end: 0.01,
|
||
epsilon_decay: 0.995,
|
||
replay_buffer_capacity: 10000,
|
||
batch_size: 32,
|
||
min_replay_size: 100,
|
||
target_update_freq: 1000,
|
||
use_double_dqn: true,
|
||
use_huber_loss: true,
|
||
huber_delta: 1.0,
|
||
leaky_relu_alpha: 0.01,
|
||
};
|
||
|
||
let mut dqn = WorkingDQN::new(config)?;
|
||
|
||
// Add experiences to replay buffer
|
||
for i in 0..200 {
|
||
let state = vec![0.1; 54];
|
||
let action = (i % 3) as u8;
|
||
let reward = (i as f32) * 0.01;
|
||
let next_state = vec![0.2; 54];
|
||
let done = i == 199;
|
||
|
||
let experience = ml::dqn::Experience::new(state, action, reward, next_state, done);
|
||
dqn.store_experience(experience)?;
|
||
}
|
||
|
||
// Run 10 training steps and verify gradients don't collapse
|
||
let mut grad_norms = Vec::new();
|
||
for _ in 0..10 {
|
||
let (loss, grad_norm) = dqn.train_step(None)?;
|
||
grad_norms.push(grad_norm);
|
||
|
||
assert!(
|
||
grad_norm > 0.0,
|
||
"Gradient norm should be positive, got {}",
|
||
grad_norm
|
||
);
|
||
assert!(loss.is_finite(), "Loss should be finite, got {}", loss);
|
||
}
|
||
|
||
// Verify gradient norms are stable (not collapsing to zero)
|
||
let avg_grad_norm = grad_norms.iter().sum::<f32>() / grad_norms.len() as f32;
|
||
assert!(
|
||
avg_grad_norm > 0.1,
|
||
"Average gradient norm too small (collapse detected): {:.6}",
|
||
avg_grad_norm
|
||
);
|
||
|
||
println!("✓ Gradient collapse prevention test passed");
|
||
println!(" - Average gradient norm: {:.4}", avg_grad_norm);
|
||
println!(
|
||
" - Min gradient norm: {:.4}",
|
||
grad_norms.iter().cloned().fold(f32::INFINITY, f32::min)
|
||
);
|
||
println!(
|
||
" - Max gradient norm: {:.4}",
|
||
grad_norms.iter().cloned().fold(f32::NEG_INFINITY, f32::max)
|
||
);
|
||
|
||
Ok(())
|
||
}
|