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

187 lines
5.8 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.
//! 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(())
}