feat(dqn): IQN+CQL integration test and verified module re-exports
Add 3 integration tests verifying the complete IQN+CQL training pipeline: full training loop with both features, IQN-only mode, and CVaR risk-aware action selection. Module re-exports for QuantileConfig/QuantileNetwork were already present from Wave 26. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
119
ml/tests/dqn_iqn_integration_test.rs
Normal file
119
ml/tests/dqn_iqn_integration_test.rs
Normal file
@@ -0,0 +1,119 @@
|
||||
//! Integration test: DQN with IQN distributional RL + CQL offline regularization
|
||||
//!
|
||||
//! Verifies the complete training loop with 2026 modernization features:
|
||||
//! - IQN replaces broken C51 (no scatter_add needed)
|
||||
//! - CQL provides offline RL regularization
|
||||
//! - CVaR enables risk-aware action selection
|
||||
|
||||
use ml::dqn::{DQNConfig, DQN, Experience};
|
||||
|
||||
#[test]
|
||||
fn test_full_iqn_cql_training_loop() {
|
||||
// Configure DQN with IQN + CQL (2026 modernization)
|
||||
let mut config = DQNConfig::default();
|
||||
config.state_dim = 8;
|
||||
config.num_actions = 3;
|
||||
config.hidden_dims = vec![32, 16];
|
||||
config.use_iqn = true;
|
||||
config.iqn_num_quantiles = 16;
|
||||
config.use_cql = true;
|
||||
config.cql_alpha = 1.0;
|
||||
config.use_distributional = false;
|
||||
config.use_dueling = false;
|
||||
config.use_per = false;
|
||||
config.batch_size = 8;
|
||||
config.min_replay_size = 8;
|
||||
config.warmup_steps = 0;
|
||||
config.epsilon_start = 0.5;
|
||||
config.use_noisy_nets = false;
|
||||
|
||||
let mut dqn = DQN::new(config).unwrap();
|
||||
|
||||
// Collect experiences via action selection
|
||||
for i in 0..20 {
|
||||
let state: Vec<f32> = (0..8).map(|j| (i * 8 + j) as f32 / 160.0).collect();
|
||||
let action = dqn.select_action(&state).unwrap();
|
||||
let reward = if i % 2 == 0 { 1.0 } else { -0.5 };
|
||||
let next_state: Vec<f32> = (0..8).map(|j| ((i + 1) * 8 + j) as f32 / 160.0).collect();
|
||||
|
||||
let exp = Experience::new(
|
||||
state,
|
||||
action.to_index() as u8,
|
||||
reward,
|
||||
next_state,
|
||||
i == 19,
|
||||
);
|
||||
dqn.store_experience(exp).unwrap();
|
||||
}
|
||||
|
||||
// Run 5 training steps
|
||||
let mut losses = Vec::new();
|
||||
for _ in 0..5 {
|
||||
let result = dqn.train_step(None);
|
||||
assert!(result.is_ok(), "Training step failed: {:?}", result.err());
|
||||
let (loss, grad_norm) = result.unwrap();
|
||||
assert!(loss.is_finite(), "Loss is not finite: {}", loss);
|
||||
assert!(grad_norm.is_finite(), "Grad norm is not finite: {}", grad_norm);
|
||||
losses.push(loss);
|
||||
}
|
||||
|
||||
// Verify loss is non-zero (model is actually learning)
|
||||
assert!(losses.iter().any(|l| *l > 0.0), "All losses are zero — model not learning");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_iqn_only_no_cql() {
|
||||
let mut config = DQNConfig::default();
|
||||
config.state_dim = 8;
|
||||
config.num_actions = 3;
|
||||
config.hidden_dims = vec![16, 16];
|
||||
config.use_iqn = true;
|
||||
config.iqn_num_quantiles = 8;
|
||||
config.use_cql = false;
|
||||
config.use_distributional = false;
|
||||
config.use_dueling = false;
|
||||
config.batch_size = 4;
|
||||
config.min_replay_size = 4;
|
||||
config.warmup_steps = 0;
|
||||
config.use_noisy_nets = false;
|
||||
|
||||
let mut dqn = DQN::new(config).unwrap();
|
||||
|
||||
for i in 0..10 {
|
||||
let exp = Experience::new(
|
||||
vec![0.1 * i as f32; 8],
|
||||
(i % 3) as u8,
|
||||
0.5,
|
||||
vec![0.2 * i as f32; 8],
|
||||
false,
|
||||
);
|
||||
dqn.store_experience(exp).unwrap();
|
||||
}
|
||||
|
||||
let result = dqn.train_step(None);
|
||||
assert!(result.is_ok(), "IQN-only training should succeed: {:?}", result.err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cvar_action_selection_integration() {
|
||||
let mut config = DQNConfig::default();
|
||||
config.state_dim = 8;
|
||||
config.num_actions = 3;
|
||||
config.hidden_dims = vec![16, 16];
|
||||
config.use_iqn = true;
|
||||
config.iqn_num_quantiles = 8;
|
||||
config.use_distributional = false;
|
||||
config.use_dueling = false;
|
||||
config.epsilon_start = 0.0;
|
||||
config.use_noisy_nets = false;
|
||||
config.warmup_steps = 0;
|
||||
config.use_cvar_action_selection = true;
|
||||
config.cvar_alpha = 0.05;
|
||||
|
||||
let mut dqn = DQN::new(config).unwrap();
|
||||
|
||||
// CVaR action selection should select more conservatively
|
||||
let state = vec![0.5f32; 8];
|
||||
let action = dqn.select_action(&state);
|
||||
assert!(action.is_ok(), "CVaR action selection should work: {:?}", action.err());
|
||||
}
|
||||
Reference in New Issue
Block a user