Detailed TDD implementation plan for: - Task 1: Bug fixes (NaN panic, Adam eps, dedup enum) - Task 2: State dimension consolidation (→51) - Task 3-5: CQL offline RL regularization - Task 4,6-9: IQN distributional RL (replaces broken C51) - Task 8: CVaR risk-aware action selection - Task 10-11: Integration test and verification Each task has exact file paths, code, test commands, and safety gates. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
45 KiB
DQN Algorithm Fix & 2026 Modernization Implementation Plan
For Claude: REQUIRED SUB-SKILL: Use superpowers:executing-plans to implement this plan task-by-task.
Goal: Fix the broken distributional RL (replace C51 with IQN), add CQL offline regularization, consolidate state dimensions, and fix training-crashing bugs.
Architecture: Add IQN (Implicit Quantile Networks) as a third loss path in DQN::train_step() alongside existing C51 and scalar paths. Add CQL regularization as a loss modifier. Wire CVaR into action selection. All changes use standard tensor ops — no scatter_add.
Tech Stack: Rust, candle-core v0.9.1 (pinned), candle-nn, candle-optimisers (vendored)
Safety gate (run after every task):
SQLX_OFFLINE=true cargo check --workspace # Must: zero errors
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::quantile_regression 2>&1 | tail -5 # Must: all pass
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -5 # Must: all pass
Task 1: Bug Fixes — NaN Panic, Adam Epsilon, Duplicate Enum
Quick wins that fix immediate crashes and paper compliance.
Files:
- Modify:
ml/src/dqn/dqn.rs:1716(NaN sort panic) - Modify:
ml/src/dqn/rainbow_agent_impl.rs:81(Adam epsilon) - Modify:
ml/src/dqn/quantile_regression.rs:312-326(remove duplicate enum)
Step 1: Fix NaN panic in dqn.rs
In ml/src/dqn/dqn.rs, find line 1716:
indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
Replace with:
indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
Step 2: Fix Adam epsilon in rainbow_agent_impl.rs
In ml/src/dqn/rainbow_agent_impl.rs, find line 81:
eps: 1e-8,
Replace with:
eps: 1.5e-4, // Rainbow paper (Hessel et al. 2018) standard for distributional stability
Step 3: Remove duplicate DistributionalType enum
In ml/src/dqn/quantile_regression.rs, delete lines 312-326 (the duplicate DistributionalType enum, Default impl, and the test_distributional_type_default test). Replace with an import:
// Re-export from distributional module (single source of truth)
pub use super::distributional::DistributionalType;
Step 4: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::quantile_regression 2>&1 | tail -5
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -5
Expected: Zero errors, all tests pass (the removed test was redundant — distributional module has its own test).
Step 5: Commit
git add ml/src/dqn/dqn.rs ml/src/dqn/rainbow_agent_impl.rs ml/src/dqn/quantile_regression.rs
git commit -m "fix(dqn): NaN sort panic, Adam eps per Rainbow paper, dedup DistributionalType"
Task 2: State Dimension Consolidation
Fix inconsistent state_dim defaults to prevent silent tensor mismatches.
Files:
- Modify:
ml/src/dqn/dqn.rs:151(DQNConfig default state_dim: 54 → 51) - Modify:
ml/src/dqn/agent.rs:191(DQNConfig default state_dim: 52 → 51)
Step 1: Fix dqn::dqn::DQNConfig::default()
In ml/src/dqn/dqn.rs, find line 151:
state_dim: 54, // Standard feature dimension
Replace with:
state_dim: 51, // 45 market features + 6 portfolio features (Wave 23)
Step 2: Fix dqn::agent::DQNConfig::default()
In ml/src/dqn/agent.rs, find line 191:
state_dim: 52, // 4 prices + 16 technical + 16 microstructure + 16 portfolio = 52
Replace with:
state_dim: 51, // 45 market features + 6 portfolio features (Wave 23)
Step 3: Update test that uses state_dim 52
In ml/src/dqn/dqn.rs, find line 2351:
config.state_dim = 52;
Replace with:
config.state_dim = 51; // Match production feature vector size (45 market + 6 portfolio)
Step 4: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -5
Expected: Zero errors, all tests pass.
Step 5: Commit
git add ml/src/dqn/dqn.rs ml/src/dqn/agent.rs
git commit -m "fix(dqn): consolidate state_dim to 51 (45 market + 6 portfolio)"
Task 3: Add CQL Config Fields to DQNConfig
Add the Conservative Q-Learning configuration parameters.
Files:
- Modify:
ml/src/dqn/dqn.rs(DQNConfig struct + Default impl)
Step 1: Add CQL fields to DQNConfig struct
In ml/src/dqn/dqn.rs, find line 145 (end of struct before closing }):
pub gradient_collapse_patience: usize,
}
Insert BEFORE the closing }:
pub gradient_collapse_patience: usize,
// Conservative Q-Learning (CQL) for offline RL (Kumar et al. 2020)
/// Whether to use CQL regularization for offline training
/// CRITICAL: Must be true when training from historical market data
/// CQL prevents Q-value overestimation on out-of-distribution actions
pub use_cql: bool,
/// CQL regularization strength (alpha parameter)
/// Higher = more conservative (lower Q-values for unseen actions)
/// Paper default: 1.0, tunable via hyperopt in range [0.1, 10.0]
pub cql_alpha: f64,
// IQN (Implicit Quantile Networks) for distributional RL (Dabney et al. 2018b)
/// Whether to use IQN instead of C51 for distributional RL
/// IQN uses quantile Huber loss (no scatter_add) — avoids Candle BUG #36
pub use_iqn: bool,
/// Number of quantile samples for IQN (more = better tail resolution)
/// Paper default: 64, trading recommendation: 200 for risk modeling
pub iqn_num_quantiles: usize,
/// Kappa parameter for quantile Huber loss
/// Controls L1/L2 transition: 1.0 = standard, higher = smoother
pub iqn_kappa: f32,
/// Cosine embedding dimension for quantile encoding
pub iqn_embedding_dim: usize,
// CVaR (Conditional Value at Risk) for risk-aware action selection
/// Whether to use CVaR-based action selection (requires use_iqn=true)
/// Selects actions that maximize worst-case returns instead of average returns
pub use_cvar_action_selection: bool,
/// CVaR confidence level: fraction of worst outcomes to optimize for
/// 0.05 = optimize for worst 5% of outcomes (conservative)
/// 0.25 = optimize for worst 25% of outcomes (moderate)
pub cvar_alpha: f32,
}
Step 2: Add defaults in Default impl
In ml/src/dqn/dqn.rs, find line 190 (end of Default impl gradient_collapse_patience):
gradient_collapse_multiplier: 2.0,
gradient_collapse_patience: 100,
}
Add CQL/IQN/CVaR defaults before closing }:
gradient_collapse_multiplier: 2.0,
gradient_collapse_patience: 100,
// CQL: Enabled by default for offline training
use_cql: true,
cql_alpha: 1.0,
// IQN: Enabled by default (replaces broken C51)
use_iqn: false, // Default false to not break existing users, enable explicitly
iqn_num_quantiles: 64,
iqn_kappa: 1.0,
iqn_embedding_dim: 64,
// CVaR: Disabled by default
use_cvar_action_selection: false,
cvar_alpha: 0.05,
}
Step 3: Add same fields to aggressive(), conservative(), emergency_safe_defaults()
For aggressive() (around line 259):
use_cql: true,
cql_alpha: 0.5, // Less conservative for aggressive exploration
use_iqn: true, // Enable IQN for distributional RL
iqn_num_quantiles: 200, // High resolution for risk
iqn_kappa: 1.0,
iqn_embedding_dim: 64,
use_cvar_action_selection: false,
cvar_alpha: 0.05,
For conservative() (around line 311):
use_cql: true,
cql_alpha: 2.0, // More conservative
use_iqn: true,
iqn_num_quantiles: 200,
iqn_kappa: 1.0,
iqn_embedding_dim: 64,
use_cvar_action_selection: true, // Risk-aware for conservative
cvar_alpha: 0.05,
For emergency_safe_defaults() (around line 372):
use_cql: true,
cql_alpha: 5.0, // Very conservative for emergency
use_iqn: false, // Disabled in emergency (simplest path)
iqn_num_quantiles: 64,
iqn_kappa: 1.0,
iqn_embedding_dim: 64,
use_cvar_action_selection: false,
cvar_alpha: 0.05,
Step 4: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -5
Expected: Zero errors, all tests pass.
Step 5: Commit
git add ml/src/dqn/dqn.rs
git commit -m "feat(dqn): add CQL, IQN, and CVaR config fields to DQNConfig"
Task 4: Expand QuantileNetwork for Multi-Action Output
Upgrade QuantileNetwork from single-action to multi-action output, add random τ sampling.
Files:
- Modify:
ml/src/dqn/quantile_regression.rs
Step 1: Add num_actions to QuantileConfig
In ml/src/dqn/quantile_regression.rs, find QuantileConfig struct (line 27):
pub struct QuantileConfig {
pub num_quantiles: usize,
pub quantile_embedding_dim: usize,
pub kappa: f32,
}
Add num_actions field:
pub struct QuantileConfig {
pub num_quantiles: usize,
pub quantile_embedding_dim: usize,
pub kappa: f32,
/// Number of actions in the DQN action space
pub num_actions: usize,
}
Update Default:
impl Default for QuantileConfig {
fn default() -> Self {
Self {
num_quantiles: 200,
quantile_embedding_dim: 64,
kappa: 1.0,
num_actions: 45, // FactoredAction space (5 exposure × 3 order × 3 urgency)
}
}
}
Step 2: Change output_layer to multi-action
In QuantileNetwork::new(), find line 86-91:
let output_layer = candle_nn::linear(
state_dim,
1, // Output single quantile value
vb.pp("quantile_output"),
)
Replace with:
let output_layer = candle_nn::linear(
state_dim,
config.num_actions, // Output Q-value per action per quantile
vb.pp("quantile_output"),
)
Step 3: Update forward() for multi-action output
Replace the forward() method (lines 108-138):
/// Forward pass: Compute quantile values Z(s, a, τ) for all actions
///
/// # Arguments
/// * `state_embed` - State embedding from base Q-network [batch, state_dim]
/// * `taus` - Quantile fractions τ ∈ [0,1] [batch, num_quantiles]
///
/// # Returns
/// Quantile values [batch, num_actions, num_quantiles]
pub fn forward(&self, state_embed: &Tensor, taus: &Tensor) -> Result<Tensor, MLError> {
let batch_size = state_embed.dim(0)?;
let num_quantiles = taus.dim(1)?;
let num_actions = self.config.num_actions;
// 1. Compute cosine embedding: ψ(τ) = [cos(πi·τ) for i in 1..embedding_dim]
let cos_embed = self.cosine_embedding(taus)?; // [batch, num_quantiles, embedding_dim]
// 2. State embedding broadcast to match quantile dimension
// [batch, state_dim] → [batch, 1, state_dim] → [batch, num_quantiles, state_dim]
let state_broadcast = state_embed
.unsqueeze(1)?
.broadcast_as((batch_size, num_quantiles, state_embed.dim(1)?))?;
// 3. Apply linear transformation to cosine embedding
// [batch, num_quantiles, embedding_dim] → [batch, num_quantiles, state_dim]
let quantile_features = self.quantile_embedding.forward(&cos_embed)
.map_err(|e| MLError::ModelError(format!("Quantile embedding forward failed: {}", e)))?;
// 4. Element-wise product: φ(s) ⊙ ψ(τ)
let combined = state_broadcast.mul(&quantile_features)?;
// 5. ReLU activation
let activated = combined.relu()?;
// 6. Project to quantile values for all actions
// [batch, num_quantiles, state_dim] → [batch, num_quantiles, num_actions]
let quantile_values = self.output_layer.forward(&activated)
.map_err(|e| MLError::ModelError(format!("Output layer forward failed: {}", e)))?;
// 7. Transpose to [batch, num_actions, num_quantiles]
quantile_values
.reshape((batch_size, num_quantiles, num_actions))?
.transpose(1, 2)
.map_err(|e| MLError::ModelError(format!("Transpose failed: {}", e)))
}
Step 4: Add random τ sampling method
After sample_uniform_quantiles, add:
/// Sample random quantiles from Uniform(0, 1) — IQN training mode
///
/// Unlike fixed quantiles (QR-DQN), IQN samples τ randomly each forward pass.
/// This enables learning a continuous quantile function.
///
/// # Arguments
/// * `batch_size` - Batch size
/// * `device` - Device to create tensor on
pub fn sample_random_quantiles(&self, batch_size: usize, device: &Device) -> CandleResult<Tensor> {
let num_quantiles = self.config.num_quantiles;
Tensor::rand(0f32, 1f32, (batch_size, num_quantiles), device)
}
Step 5: Update to_scalar and compute_cvar for multi-action
Replace to_scalar (line 216-219):
/// Compute expected Q-values from quantile distributions (mean over quantiles)
///
/// # Arguments
/// * `quantiles` - Quantile values [batch, num_actions, num_quantiles]
///
/// # Returns
/// Expected Q-values [batch, num_actions]
pub fn to_expected_q(&self, quantiles: &Tensor) -> CandleResult<Tensor> {
quantiles.mean(2) // Average over quantiles dimension
}
Replace compute_cvar (lines 236-246):
/// Extract CVaR (Conditional Value at Risk) for each action
///
/// CVaR_α = E[Z | Z ≤ VaR_α] = mean of bottom α quantiles
///
/// # Arguments
/// * `quantiles` - Quantile values [batch, num_actions, num_quantiles]
/// * `alpha` - Risk level (e.g., 0.05 for worst 5%)
///
/// # Returns
/// CVaR values per action [batch, num_actions]
pub fn compute_cvar(&self, quantiles: &Tensor, alpha: f32) -> CandleResult<Tensor> {
let num_quantiles = self.config.num_quantiles;
let num_tail = (num_quantiles as f32 * alpha).ceil() as usize;
let num_tail = num_tail.max(1);
// Extract bottom α quantiles (assumes quantiles are ordered)
let tail = quantiles.narrow(2, 0, num_tail)?;
tail.mean(2)
}
Step 6: Fix existing tests for new signature
Update all tests that create QuantileConfig to include num_actions:
// In every test that uses QuantileConfig::default(), no change needed
// (default already has num_actions: 45)
// In tests that create explicit QuantileConfig:
let config = QuantileConfig {
num_quantiles: 200,
quantile_embedding_dim: 64,
kappa: 1.0,
num_actions: 3, // Small action space for fast tests
};
Update test_quantile_network_forward to check new shape:
assert_eq!(quantile_values.shape().dims(), &[batch_size, config.num_actions, config.num_quantiles]);
Update test_quantile_to_scalar → test_to_expected_q:
#[test]
fn test_to_expected_q() -> Result<(), MLError> {
let config = QuantileConfig { num_actions: 3, ..Default::default() };
let device = Device::Cpu;
let vb = VarBuilder::zeros(DType::F32, &device);
let network = QuantileNetwork::new(&config, 64, vb)?;
let batch_size = 4;
let quantiles = Tensor::randn(0f32, 1f32, (batch_size, config.num_actions, config.num_quantiles), &device)
.map_err(|e| MLError::ModelError(e.to_string()))?;
let expected_q = network.to_expected_q(&quantiles)
.map_err(|e| MLError::ModelError(e.to_string()))?;
assert_eq!(expected_q.shape().dims(), &[batch_size, config.num_actions]);
Ok(())
}
Update test_cvar_computation for new shape:
#[test]
fn test_cvar_computation() -> Result<(), MLError> {
let config = QuantileConfig { num_actions: 3, ..Default::default() };
let device = Device::Cpu;
let vb = VarBuilder::zeros(DType::F32, &device);
let network = QuantileNetwork::new(&config, 64, vb)?;
let batch_size = 4;
// Create ascending quantile values [batch, num_actions, num_quantiles]
let quantiles = Tensor::arange(0f32, (batch_size * config.num_actions * config.num_quantiles) as f32, &device)
.map_err(|e| MLError::ModelError(e.to_string()))?
.reshape((batch_size, config.num_actions, config.num_quantiles))
.map_err(|e| MLError::ModelError(e.to_string()))?;
let alpha = 0.05;
let cvar = network.compute_cvar(&quantiles, alpha)
.map_err(|e| MLError::ModelError(e.to_string()))?;
assert_eq!(cvar.shape().dims(), &[batch_size, config.num_actions]);
// CVaR should be lower than mean for ascending quantiles
let mean_val = network.to_expected_q(&quantiles)
.map_err(|e| MLError::ModelError(e.to_string()))?;
let cvar_vec: Vec<f32> = cvar.flatten_all()
.map_err(|e| MLError::ModelError(e.to_string()))?
.to_vec1()
.map_err(|e| MLError::ModelError(e.to_string()))?;
let mean_vec: Vec<f32> = mean_val.flatten_all()
.map_err(|e| MLError::ModelError(e.to_string()))?
.to_vec1()
.map_err(|e| MLError::ModelError(e.to_string()))?;
for (c, m) in cvar_vec.iter().zip(mean_vec.iter()) {
assert!(c < m, "CVaR should be less than mean for ascending quantiles");
}
Ok(())
}
Step 7: Add test for random tau sampling
#[test]
fn test_random_quantile_sampling() -> Result<(), MLError> {
let config = QuantileConfig { num_actions: 3, ..Default::default() };
let device = Device::Cpu;
let vb = VarBuilder::zeros(DType::F32, &device);
let network = QuantileNetwork::new(&config, 64, vb)?;
let batch_size = 4;
let taus = network.sample_random_quantiles(batch_size, &device)
.map_err(|e| MLError::ModelError(e.to_string()))?;
assert_eq!(taus.shape().dims(), &[batch_size, config.num_quantiles]);
// All values should be in (0, 1)
let taus_vec: Vec<f32> = taus.flatten_all()
.map_err(|e| MLError::ModelError(e.to_string()))?
.to_vec1()
.map_err(|e| MLError::ModelError(e.to_string()))?;
for tau in &taus_vec {
assert!(*tau >= 0.0 && *tau <= 1.0, "Tau {} out of range", tau);
}
Ok(())
}
Step 8: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::quantile_regression 2>&1 | tail -15
Expected: Zero errors, all tests pass.
Step 9: Commit
git add ml/src/dqn/quantile_regression.rs
git commit -m "feat(dqn): expand QuantileNetwork to multi-action IQN with random tau sampling"
Task 5: Add CQL Loss Computation to train_step()
Implement the Conservative Q-Learning regularization in the training loop.
Files:
- Modify:
ml/src/dqn/dqn.rs(train_step method)
Step 1: Write CQL test
In ml/src/dqn/dqn.rs, in the #[cfg(test)] mod tests block, add:
#[test]
fn test_cql_regularization() {
let mut config = DQNConfig::default();
config.state_dim = 8;
config.num_actions = 3;
config.hidden_dims = vec![16, 16];
config.use_cql = true;
config.cql_alpha = 1.0;
config.use_iqn = false;
config.use_distributional = false;
config.use_dueling = false;
config.batch_size = 4;
config.min_replay_size = 2;
let mut dqn = DQN::new(config).unwrap();
// Add enough experiences
for i in 0..10 {
let exp = Experience {
state: vec![0.1; 8],
action: (i % 3) as usize,
reward: crate::dqn::experience::Reward::Simple(0.5),
next_state: vec![0.2; 8],
done: false,
};
dqn.store_experience(exp);
}
// Train step should succeed with CQL enabled
let result = dqn.train_step(None);
assert!(result.is_ok(), "Training with CQL should succeed: {:?}", result.err());
let (loss, _grad_norm) = result.unwrap();
assert!(loss.is_finite(), "CQL loss should be finite, got {}", loss);
}
Step 2: Run test to verify it fails
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests::test_cql_regularization -v 2>&1 | tail -10
Expected: Compile error (new config fields don't exist yet) or test compiles but CQL code path doesn't exist.
Step 3: Implement CQL in train_step()
In ml/src/dqn/dqn.rs, find the entropy regularization section (around line 1580):
// Add entropy regularization (in-training diversity penalty)
let entropy_penalty = self.calculate_entropy_penalty()?;
let entropy_weight = 0.1;
let entropy_term = (entropy_penalty * entropy_weight)?;
let loss = loss_value.add(&entropy_term)?;
Insert CQL computation AFTER entropy and BEFORE extracting loss scalar:
// Add entropy regularization (in-training diversity penalty)
let entropy_penalty = self.calculate_entropy_penalty()?;
let entropy_weight = 0.1;
let entropy_term = (entropy_penalty * entropy_weight)?;
let loss_with_entropy = loss_value.add(&entropy_term)?;
// CQL regularization for offline RL (Kumar et al. 2020)
// Penalizes high Q-values for actions not taken in the training data
// CQL_penalty = E[logsumexp(Q(s, all_actions))] - E[Q(s, a_data)]
let loss = if self.config.use_cql {
// logsumexp(Q(s, all_actions)) for numerical stability:
// logsumexp(x) = max(x) + log(sum(exp(x - max(x))))
let q_max = current_q_values.max(1)?; // [batch]
let q_max_broadcast = q_max.unsqueeze(1)?
.broadcast_as(current_q_values.shape())?;
let q_shifted = (current_q_values.detach() - q_max_broadcast)?;
let logsumexp = (q_shifted.exp()?.sum(1)?.log()? + q_max)?; // [batch]
// Q(s, a_data) — Q-values for actions actually taken in the dataset
let q_data = state_action_values.detach(); // [batch], already gathered above
// CQL penalty = mean(logsumexp - Q_data)
let cql_penalty = (logsumexp - q_data)?.mean_all()?;
// Total loss = td_loss + entropy + alpha * cql_penalty
let cql_term = (cql_penalty * self.config.cql_alpha)?;
if self.training_steps % 100 == 0 {
let cql_val: f32 = cql_term.to_scalar()
.unwrap_or(0.0);
tracing::info!(
"CQL penalty at step {}: {:.4} (alpha: {:.2})",
self.training_steps, cql_val, self.config.cql_alpha
);
}
loss_with_entropy.add(&cql_term)?
} else {
loss_with_entropy
};
Step 4: Run test
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests::test_cql_regularization -v 2>&1 | tail -10
Expected: PASS
Step 5: Run full safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -10
Expected: Zero errors, ALL DQN tests pass.
Step 6: Commit
git add ml/src/dqn/dqn.rs
git commit -m "feat(dqn): add CQL offline RL regularization to train_step (Kumar et al. 2020)"
Task 6: Add IQN Fields to DQN Struct and Initialization
Add the IQN network and target network to the DQN struct, initialize in DQN::new().
Files:
- Modify:
ml/src/dqn/dqn.rs(DQN struct, DQN::new())
Step 1: Add IQN fields to DQN struct
In ml/src/dqn/dqn.rs, find the DQN struct (line ~641). Add after the nstep_buffer field:
/// IQN (Implicit Quantile Network) for distributional RL (Dabney et al. 2018b)
/// Replaces C51's scatter_add-based projection with quantile Huber loss
iqn_network: Option<super::quantile_regression::QuantileNetwork>,
/// IQN target network (frozen copy for stable target computation)
iqn_target_network: Option<super::quantile_regression::QuantileNetwork>,
Step 2: Initialize IQN in DQN::new()
In DQN::new(), find the N-step buffer creation (around line 807). After it, add:
// IQN initialization (when use_iqn=true, replaces C51)
let (iqn_network, iqn_target_network) = if config.use_iqn {
let iqn_config = super::quantile_regression::QuantileConfig {
num_quantiles: config.iqn_num_quantiles,
quantile_embedding_dim: config.iqn_embedding_dim,
kappa: config.iqn_kappa,
num_actions: config.num_actions,
};
// Get state embedding dimension from the last hidden layer
let embed_dim = config.hidden_dims.last().copied().unwrap_or(config.state_dim);
// Create IQN main network
let iqn_vars = VarMap::new();
let iqn_vb = VarBuilder::from_varmap(&iqn_vars, DType::F32, &device);
let iqn_net = super::quantile_regression::QuantileNetwork::new(
&iqn_config, embed_dim, iqn_vb
)?;
// Create IQN target network
let iqn_target_vars = VarMap::new();
let iqn_target_vb = VarBuilder::from_varmap(&iqn_target_vars, DType::F32, &device);
let iqn_target = super::quantile_regression::QuantileNetwork::new(
&iqn_config, embed_dim, iqn_target_vb
)?;
(Some(iqn_net), Some(iqn_target))
} else {
(None, None)
};
Step 3: Add IQN fields to Self initialization
In the Ok(Self { ... }) block, add:
iqn_network,
iqn_target_network,
Step 4: Write test for IQN initialization
#[test]
fn test_dqn_with_iqn_creation() {
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 = 32;
config.use_distributional = false;
config.use_dueling = false;
let dqn = DQN::new(config);
assert!(dqn.is_ok(), "DQN with IQN should create successfully: {:?}", dqn.err());
}
Step 5: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -10
Expected: Zero errors, all tests pass.
Step 6: Commit
git add ml/src/dqn/dqn.rs
git commit -m "feat(dqn): add IQN network fields to DQN struct with initialization"
Task 7: IQN Loss Path in train_step()
Add the IQN quantile Huber loss computation as a third loss path.
Files:
- Modify:
ml/src/dqn/dqn.rs(train_step method)
Step 1: Write IQN training test
#[test]
fn test_iqn_training_step() {
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 = 2;
let mut dqn = DQN::new(config).unwrap();
// Add experiences
for i in 0..10 {
let exp = Experience {
state: vec![0.1 * i as f32; 8],
action: (i % 3) as usize,
reward: crate::dqn::experience::Reward::Simple(0.5),
next_state: vec![0.2 * i as f32; 8],
done: i == 9,
};
dqn.store_experience(exp);
}
// Train step with IQN should succeed
let result = dqn.train_step(None);
assert!(result.is_ok(), "IQN training step should succeed: {:?}", result.err());
let (loss, grad_norm) = result.unwrap();
assert!(loss.is_finite(), "IQN loss should be finite: {}", loss);
assert!(grad_norm >= 0.0, "Gradient norm should be non-negative: {}", grad_norm);
}
Step 2: Add IQN loss path in train_step()
In train_step(), find the loss computation section (around line 1377 where C51 path starts). The structure is:
let loss_value = if self.dist_dueling_q_network.is_some() {
// C51 CATEGORICAL LOSS PATH ...
} else {
// STANDARD DQN SCALAR LOSS PATH ...
};
Change this to a 3-way branch — IQN first (highest priority for new code):
let loss_value = if self.config.use_iqn && self.iqn_network.is_some() {
// IQN QUANTILE HUBER LOSS PATH (Dabney et al. 2018b)
// Uses quantile Huber loss — no scatter_add, no gradient flow issues
let iqn_net = self.iqn_network.as_ref().unwrap();
let iqn_target = self.iqn_target_network.as_ref().unwrap();
// Get state embeddings from the base Q-network's hidden layers
// Forward through all hidden layers except the output layer
let state_embed = self.get_state_embedding(&states_tensor)?;
let next_state_embed = self.get_state_embedding(&next_states_tensor)?;
// Sample random quantiles for training (IQN: τ ~ Uniform(0,1))
let taus = iqn_net.sample_random_quantiles(batch_size, device)
.map_err(|e| MLError::TrainingError(format!("Failed to sample taus: {}", e)))?;
let target_taus = iqn_target.sample_random_quantiles(batch_size, device)
.map_err(|e| MLError::TrainingError(format!("Failed to sample target taus: {}", e)))?;
// Forward through IQN: [batch, num_actions, num_quantiles]
let all_quantiles = iqn_net.forward(&state_embed, &taus)?;
// Gather quantiles for taken actions: [batch, 1, num_quantiles] → [batch, num_quantiles]
let actions_for_gather = actions_tensor
.unsqueeze(1)? // [batch, 1]
.unsqueeze(2)? // [batch, 1, 1]
.broadcast_as((batch_size, 1, self.config.iqn_num_quantiles))?
.contiguous()?;
let current_quantiles = all_quantiles.contiguous()?
.gather(&actions_for_gather, 1)?
.squeeze(1)?; // [batch, num_quantiles]
// Target quantiles (detached — no gradient flow)
let next_all_quantiles = iqn_target.forward(&next_state_embed.detach(), &target_taus)
.map_err(|e| MLError::TrainingError(format!("IQN target forward failed: {}", e)))?
.detach();
// Select best next actions using online network (Double DQN)
let next_q_values = iqn_net.to_expected_q(&all_quantiles.detach())
.map_err(|e| MLError::TrainingError(format!("IQN to_expected_q failed: {}", e)))?;
let next_actions = next_q_values.argmax(1)?; // [batch]
// Gather target quantiles for best actions
let next_actions_for_gather = next_actions
.unsqueeze(1)?
.unsqueeze(2)?
.broadcast_as((batch_size, 1, self.config.iqn_num_quantiles))?
.contiguous()?;
let next_quantiles = next_all_quantiles.contiguous()?
.gather(&next_actions_for_gather, 1)?
.squeeze(1)?; // [batch, num_quantiles]
// Compute target quantiles: r + γ * (1 - done) * Z_target
let gamma = self.config.gamma;
let rewards_broadcast = rewards_tensor
.unsqueeze(1)?
.broadcast_as((batch_size, self.config.iqn_num_quantiles))?;
let not_done_broadcast = (Tensor::ones(&[batch_size], DType::F32, device)? - &dones_tensor)?
.unsqueeze(1)?
.broadcast_as((batch_size, self.config.iqn_num_quantiles))?;
let target_quantiles = (rewards_broadcast + (next_quantiles * not_done_broadcast)? * gamma as f64)?
.detach();
// Quantile Huber loss
let loss = super::quantile_regression::quantile_huber_loss(
¤t_quantiles,
&target_quantiles,
&taus,
self.config.iqn_kappa,
).map_err(|e| MLError::TrainingError(format!("Quantile Huber loss failed: {}", e)))?;
// Apply PER importance sampling weights
let weights_tensor = Tensor::from_vec(weights.clone(), batch_size, device)?
.detach();
// Expand weights to match quantile dimension, then reduce
let per_sample_loss = super::quantile_regression::quantile_huber_loss(
¤t_quantiles,
&target_quantiles,
&taus,
self.config.iqn_kappa,
).map_err(|e| MLError::TrainingError(format!("QHL per-sample failed: {}", e)))?;
per_sample_loss
} else if self.dist_dueling_q_network.is_some() {
// C51 CATEGORICAL LOSS PATH (existing — kept for backwards compatibility)
// ... existing C51 code unchanged ...
Step 3: Add get_state_embedding() helper method
Add this method to the impl DQN block:
/// Get state embeddings from the base Q-network's hidden layers
/// Forwards through all hidden layers except the output, producing
/// the intermediate representation needed by IQN.
///
/// # Returns
/// State embeddings [batch, last_hidden_dim]
fn get_state_embedding(&self, states: &Tensor) -> Result<Tensor, MLError> {
let states = states.to_device(&self.device)?;
if self.q_network.use_noisy_nets {
// Use noisy layers (skip output layer)
let mut x = states;
let num_layers = self.q_network.noisy_layers.len();
for (i, layer) in self.q_network.noisy_layers.iter().enumerate() {
if i >= num_layers - 1 { break; } // Skip output layer
x = layer.forward(&x)?;
x = candle_nn::ops::leaky_relu(&x, self.q_network.leaky_relu_alpha)?;
}
Ok(x)
} else {
// Use standard layers (skip output layer)
let mut x = states;
let num_layers = self.q_network.layers.len();
for (i, layer) in self.q_network.layers.iter().enumerate() {
if i >= num_layers - 1 { break; } // Skip output layer
x = layer.forward(&x).map_err(|e| {
MLError::ModelError(format!("Embedding forward failed at layer {}: {}", i, e))
})?;
x = candle_nn::ops::leaky_relu(&x, self.q_network.leaky_relu_alpha)?;
}
Ok(x)
}
}
Step 4: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests::test_iqn 2>&1 | tail -10
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -10
Expected: Zero errors, all tests pass.
Step 5: Commit
git add ml/src/dqn/dqn.rs
git commit -m "feat(dqn): add IQN quantile Huber loss path in train_step (replaces C51)"
Task 8: IQN Action Selection + CVaR
Wire IQN into select_action() with optional CVaR risk-aware mode.
Files:
- Modify:
ml/src/dqn/dqn.rs(select_action method)
Step 1: Write CVaR action selection test
#[test]
fn test_iqn_action_selection() {
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; // Force greedy for testing
config.use_noisy_nets = false;
config.warmup_steps = 0;
let mut dqn = DQN::new(config).unwrap();
let state = vec![0.5f32; 8];
let action = dqn.select_action(&state);
assert!(action.is_ok(), "IQN action selection should succeed: {:?}", action.err());
}
Step 2: Modify select_action() for IQN
In select_action(), find the greedy action selection block (around line 1023). Replace the greedy branch:
} else {
// Greedy action selection
let state_tensor = Tensor::from_vec(
state.to_vec(),
(1, self.config.state_dim),
self.q_network.device(),
)?;
if self.config.use_iqn && self.iqn_network.is_some() {
// IQN: Compute quantile-based Q-values
let iqn_net = self.iqn_network.as_ref().unwrap();
let state_embed = self.get_state_embedding(&state_tensor)?;
// Use fixed quantiles for deterministic action selection
let taus = iqn_net.sample_uniform_quantiles(1, &self.device)
.map_err(|e| MLError::ModelError(format!("Failed to sample taus: {}", e)))?;
// Forward: [1, num_actions, num_quantiles]
let all_quantiles = iqn_net.forward(&state_embed, &taus)?;
// Score each action
let action_scores = if self.config.use_cvar_action_selection {
// CVaR: Optimize for worst-case outcomes
iqn_net.compute_cvar(&all_quantiles, self.config.cvar_alpha)
.map_err(|e| MLError::ModelError(format!("CVaR computation failed: {}", e)))?
} else {
// Standard: Mean over quantiles
iqn_net.to_expected_q(&all_quantiles)
.map_err(|e| MLError::ModelError(format!("Expected Q computation failed: {}", e)))?
};
// Argmax over actions
let best_action_idx = action_scores
.argmax(1)?
.get(0)?
.to_scalar::<u32>()?;
FactoredAction::from_index(best_action_idx as usize)?
} else {
// Standard Q-network action selection (existing code)
let q_values = self.forward(&state_tensor)?;
let best_action_idx = q_values
.argmax(1)?
.get(0)?
.to_scalar::<u32>()?;
FactoredAction::from_index(best_action_idx as usize)?
}
};
Step 3: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn::dqn::tests 2>&1 | tail -10
Expected: Zero errors, all tests pass.
Step 4: Commit
git add ml/src/dqn/dqn.rs
git commit -m "feat(dqn): add IQN action selection with optional CVaR risk-aware mode"
Task 9: IQN Target Network Updates
Wire IQN target network into the Polyak/hard update logic in train_step().
Files:
- Modify:
ml/src/dqn/dqn.rs(target update section of train_step) - Modify:
ml/src/dqn/quantile_regression.rs(add VarMap access + copy_weights)
Step 1: Add VarMap accessor and copy_weights to QuantileNetwork
In ml/src/dqn/quantile_regression.rs, after the QuantileNetwork struct, add the vars field and accessor. This requires storing the VarMap:
Add vars: VarMap field to the struct. Update new() to accept and store VarMap. Add:
/// Get network variables for optimizer
pub fn vars(&self) -> &VarMap {
&self.vars
}
/// Copy weights from another QuantileNetwork (for target network sync)
pub fn copy_weights_from(&mut self, other: &QuantileNetwork) -> Result<(), MLError> {
let self_data = self.vars.data().lock()
.map_err(|e| MLError::ConcurrencyError { operation: format!("lock self vars: {}", e) })?;
let other_data = other.vars.data().lock()
.map_err(|e| MLError::ConcurrencyError { operation: format!("lock other vars: {}", e) })?;
for (name, self_var) in self_data.iter() {
if let Some(other_var) = other_data.get(name) {
self_var.set(other_var.as_tensor())
.map_err(|e| MLError::ModelError(format!("Failed to copy weight {}: {}", name, e)))?;
}
}
Ok(())
}
Step 2: Add IQN to target update section in train_step()
In the soft update section (around line 1629), add IQN target update:
// Also update IQN target network if present
if let (Some(ref iqn_net), Some(ref iqn_target)) =
(&self.iqn_network, &self.iqn_target_network)
{
polyak_update(iqn_net.vars(), iqn_target.vars(), self.config.tau)
.map_err(|e| MLError::TrainingError(format!("IQN Polyak update failed: {}", e)))?;
}
Do the same in the hard update section.
Step 3: Run safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn 2>&1 | tail -10
Step 4: Commit
git add ml/src/dqn/dqn.rs ml/src/dqn/quantile_regression.rs
git commit -m "feat(dqn): wire IQN target network into Polyak/hard update logic"
Task 10: Update mod.rs Re-exports and Integration Test
Update module re-exports and write an end-to-end integration test.
Files:
- Modify:
ml/src/dqn/mod.rs(re-exports) - Create:
ml/tests/dqn_iqn_integration_test.rs(integration test)
Step 1: Update mod.rs re-exports
In ml/src/dqn/mod.rs, update the quantile regression re-exports:
// Re-export IQN/QR-DQN components (Wave 26 P1.13 + 2026 modernization)
pub use quantile_regression::{QuantileConfig, QuantileNetwork, quantile_huber_loss};
Step 2: Write integration test
Create ml/tests/dqn_iqn_integration_test.rs:
//! Integration test: DQN with IQN distributional RL + CQL offline regularization
//!
//! Verifies the complete training loop with 2026 modernization features.
use ml::dqn::{DQNConfig, DQN, Experience};
use ml::dqn::experience::Reward;
#[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;
let mut dqn = DQN::new(config).unwrap();
// Collect experiences
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 action.to_legacy_action() == 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 {
state,
action: action.to_index(),
reward: Reward::Simple(reward),
next_state,
done: i == 19,
};
dqn.store_experience(exp);
}
// 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");
}
Step 3: Run integration test
SQLX_OFFLINE=true cargo test -p ml --test dqn_iqn_integration_test -v 2>&1 | tail -10
Expected: PASS
Step 4: Run full safety gate
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true cargo test -p ml 2>&1 | tail -20
Expected: Zero errors, all tests pass.
Step 5: Commit
git add ml/src/dqn/mod.rs ml/tests/dqn_iqn_integration_test.rs
git commit -m "feat(dqn): IQN+CQL integration test and updated module re-exports"
Task 11: Final Verification and Documentation Update
Comprehensive verification that everything works together.
Step 1: Run all DQN-related tests
SQLX_OFFLINE=true cargo test -p ml --lib -- dqn 2>&1 | tail -30
SQLX_OFFLINE=true cargo test -p ml --test dqn_iqn_integration_test -v 2>&1 | tail -10
Step 2: Run full workspace check
SQLX_OFFLINE=true cargo check --workspace
Step 3: Verify no regressions in other tests
SQLX_OFFLINE=true cargo test -p ml --lib 2>&1 | tail -5
Step 4: Commit any final adjustments
If any test fixes were needed:
git add -A && git commit -m "fix(dqn): final adjustments for IQN+CQL integration"
Summary
| Task | What | Files | Key Risk |
|---|---|---|---|
| 1 | Bug fixes (NaN, Adam eps, dedup enum) | dqn.rs, rainbow_agent_impl.rs, quantile_regression.rs | Low |
| 2 | State dim consolidation (→51) | dqn.rs, agent.rs | Low |
| 3 | CQL config fields | dqn.rs | Low |
| 4 | QuantileNetwork multi-action + random τ | quantile_regression.rs | Medium (tensor shapes) |
| 5 | CQL loss in train_step() | dqn.rs | Medium (gradient flow) |
| 6 | IQN fields in DQN struct | dqn.rs | Medium (initialization) |
| 7 | IQN loss path in train_step() | dqn.rs | High (tensor shapes, gather ops) |
| 8 | IQN action selection + CVaR | dqn.rs | Medium |
| 9 | IQN target network updates | dqn.rs, quantile_regression.rs | Low |
| 10 | Integration test + re-exports | mod.rs, tests/ | Low |
| 11 | Final verification | all | Low |