Files
foxhunt/docs/plans/2026-02-20-dqn-algorithm-fix-implementation.md
jgrusewski 6007f98c26 docs: Add DQN algorithm fix implementation plan (11 tasks)
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>
2026-02-20 14:44:44 +01:00

45 KiB
Raw Blame History

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_scalartest_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(
                &current_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(
                &current_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