diff --git a/crates/ml/src/dqn/dqn.rs b/crates/ml/src/dqn/dqn.rs index c6ad42c65..48c24b56b 100644 --- a/crates/ml/src/dqn/dqn.rs +++ b/crates/ml/src/dqn/dqn.rs @@ -235,8 +235,8 @@ impl Default for DQNConfig { dueling_hidden_dim: 128, use_distributional: true, // BUG #36 FIXED: C51 re-enabled num_atoms: 51, - v_min: -10.0, // Widened: [-2,2] compressed Q-values, [-10,10] gives C51 room to separate actions - v_max: 10.0, // Widened: [-2,2] compressed Q-values, [-10,10] gives C51 room to separate actions + v_min: -25.0, // DSR Q-values: rewards ±2 with gamma=0.92 → Q ≈ ±25 + v_max: 25.0, // DSR Q-values: rewards ±2 with gamma=0.92 → Q ≈ ±25 use_noisy_nets: false, noisy_sigma_init: 0.5, enable_q_value_clipping: true, @@ -862,7 +862,7 @@ impl Sequential { device: Device, leaky_relu_alpha: f64, use_noisy_nets: bool, - _noisy_sigma_init: f64, // Reserved for future custom sigma initialization + noisy_sigma_init: f64, ) -> Result { let vars = VarMap::new(); let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); @@ -876,7 +876,7 @@ impl Sequential { for (i, &hidden_dim) in hidden_dims.iter().enumerate() { let layer_name = format!("noisy_hidden_{}", i); let layer_vb = var_builder.pp(&layer_name); - let noisy_layer = super::noisy_layers::NoisyLinear::new(current_dim, hidden_dim, layer_vb) + let noisy_layer = super::noisy_layers::NoisyLinear::new(current_dim, hidden_dim, layer_vb, noisy_sigma_init) .map_err(|e| MLError::ModelError(format!("Failed to create noisy layer {}: {}", i, e)))?; noisy_layers.push(noisy_layer); current_dim = hidden_dim; @@ -884,7 +884,7 @@ impl Sequential { // Output layer also noisy let output_vb = var_builder.pp("noisy_output"); - let noisy_output = super::noisy_layers::NoisyLinear::new(current_dim, output_dim, output_vb) + let noisy_output = super::noisy_layers::NoisyLinear::new(current_dim, output_dim, output_vb, noisy_sigma_init) .map_err(|e| MLError::ModelError(format!("Failed to create noisy output layer: {}", e)))?; noisy_layers.push(noisy_output); } else { @@ -1458,7 +1458,7 @@ impl DQN { // triple-stacking (noisy + epsilon + count) while ensuring some random // actions feed diverse experiences into the replay buffer. let effective_epsilon = if self.config.use_noisy_nets { - self.config.noisy_epsilon_floor // Minimum exploration floor for noisy nets + self.config.noisy_epsilon_floor.max(0.02) // Minimum 2% random exploration } else { self.epsilon // Standard epsilon-greedy exploration }; @@ -1610,7 +1610,7 @@ impl DQN { // Noisy nets + epsilon floor (same logic as select_action) let effective_epsilon = if self.config.use_noisy_nets { - self.config.noisy_epsilon_floor + self.config.noisy_epsilon_floor.max(0.02) // Minimum 2% random exploration } else { self.epsilon }; @@ -1826,11 +1826,43 @@ impl DQN { } } + /// Compute Q-values for a batch of states, dispatching to the IQN network + /// when `use_iqn=true`, otherwise falling back to `self.forward()`. + /// + /// Returns a `[N, num_actions]` tensor of expected Q-values (or CVaR values + /// when `use_cvar_action_selection` is enabled). + fn q_values_for_batch(&self, states: &Tensor) -> Result { + if self.config.use_iqn && self.iqn_network.is_some() { + let iqn_net = self.iqn_network.as_ref().ok_or_else(|| { + MLError::ModelError("IQN network not initialized despite use_iqn=true".into()) + })?; + let state_embed = self.get_state_embedding(states)?; + let batch_size = state_embed.dim(0)?; + let taus = iqn_net.sample_uniform_quantiles(batch_size, &self.device) + .map_err(|e| MLError::ModelError(format!("Failed to sample taus: {}", e)))?; + // Forward: [batch, num_actions, num_quantiles] + let all_quantiles = iqn_net.forward(&state_embed, &taus)?; + // Reduce to [batch, num_actions] + if self.config.use_cvar_action_selection { + iqn_net.compute_cvar(&all_quantiles, self.config.cvar_alpha) + .map_err(|e| MLError::ModelError(format!("CVaR computation failed: {}", e))) + } else { + iqn_net.to_expected_q(&all_quantiles) + .map_err(|e| MLError::ModelError(format!("Expected Q computation failed: {}", e))) + } + } else { + self.forward(states) + } + } + /// Batch greedy action selection — single GPU forward pass for N states. /// /// Runs a batched forward pass through the Q-network and returns the greedy /// (argmax) action index for each state. No exploration, no side effects. /// + /// When `use_iqn=true`, routes through the IQN network so that inference + /// is consistent with the training loss path. + /// /// # Arguments /// * `states` - Tensor of shape `[N, state_dim]` /// @@ -1841,7 +1873,7 @@ impl DQN { /// Reduces N individual GPU kernel launches to a single batched forward pass. /// With EVAL_CHUNK_SIZE=1024, this gives ~1000× fewer kernel launches vs per-bar inference. pub fn batch_greedy_actions(&self, states: &Tensor) -> Result, MLError> { - let q_values = self.forward(states)?; + let q_values = self.q_values_for_batch(states)?; let action_indices = q_values .argmax(1) .map_err(|e| MLError::ModelError(format!("Batch argmax failed: {}", e)))? @@ -1858,6 +1890,9 @@ impl DQN { /// CPU transfer of Q-values. Uses: argmax(Q/T - log(-log(U))) where U~Uniform(0,1). /// At low temperature (~0.1), nearly greedy; provides diversity when Q-values are uniform. /// + /// When `use_iqn=true`, routes through the IQN network so that inference + /// is consistent with the training loss path. + /// /// # Arguments /// * `states` - Tensor of shape `[N, state_dim]` /// * `temperature` - Boltzmann temperature (clamped to >= 1e-6) @@ -1866,7 +1901,7 @@ impl DQN { states: &Tensor, temperature: f64, ) -> Result, MLError> { - let q_values = self.forward(states)?; + let q_values = self.q_values_for_batch(states)?; let temp = temperature.max(1e-6); let temp_tensor = Tensor::new(&[temp as f32], &self.device) @@ -1902,12 +1937,15 @@ impl DQN { /// DQN outputs [N, 5] Q-values (one per exposure level). Gumbel-max samples /// stochastically from the Q-value distribution with temperature control. /// Only transfers N u32 indices (exposure levels 0-4) to CPU. + /// + /// When `use_iqn=true`, routes through the IQN network so that inference + /// is consistent with the training loss path. pub fn batch_hierarchical_softmax_actions( &self, states: &Tensor, temperature: f64, ) -> Result, MLError> { - let q_values = self.forward(states)?; // [N, 5] on GPU + let q_values = self.q_values_for_batch(states)?; // [N, 5] on GPU let (_n, _) = q_values.dims2().map_err(|e| { MLError::ModelError(format!("Expected 2D Q-values: {}", e)) })?; @@ -3889,6 +3927,70 @@ mod tests { Ok(()) } + /// Verify batch_greedy_actions routes through IQN when use_iqn=true, + /// and that IQN loss decreases over multiple training steps. + #[test] + fn test_iqn_batch_greedy_actions_uses_iqn_network() -> anyhow::Result<()> { + 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.use_cql = false; + config.batch_size = 4; + config.min_replay_size = 2; + config.epsilon_start = 0.0; + config.use_noisy_nets = false; + config.warmup_steps = 0; + + let mut dqn = DQN::new(config)?; + + // Fill replay buffer with structured data so training can learn + for i in 0..20 { + let exp = Experience::new( + vec![0.1 * i as f32; 8], + (i % 3) as u8, + if i % 3 == 0 { 1.0 } else { -0.5 }, + vec![0.2 * i as f32; 8], + i == 19, + ); + dqn.store_experience(exp)?; + } + + // Train multiple steps and verify loss is finite + let mut losses = Vec::new(); + for _ in 0..10 { + let (loss, grad_norm) = dqn.train_step(None)?; + assert!(loss.is_finite(), "IQN loss should be finite: {}", loss); + assert!(grad_norm >= 0.0, "Grad norm should be non-negative"); + losses.push(loss); + } + + // Verify batch_greedy_actions succeeds with IQN (this would fail before + // the fix because it would use the standard Q-network instead of IQN) + let states = Tensor::randn(0_f32, 1.0, (4, 8), &candle_core::Device::Cpu)?; + let actions = dqn.batch_greedy_actions(&states)?; + assert_eq!(actions.len(), 4, "Should return one action per state"); + for &a in &actions { + assert!(a < 3, "Action index should be < num_actions"); + } + + // Verify batch_softmax_actions also works with IQN + let softmax_actions = dqn.batch_softmax_actions(&states, 1.0)?; + assert_eq!(softmax_actions.len(), 4, "Softmax should return one action per state"); + + // Verify select_action_inference works (read-only path) + let state = vec![0.5_f32; 8]; + let (action, confidence) = dqn.select_action_inference(&state)?; + assert!((action.exposure as usize) < 5, "Exposure index should be valid"); + assert!(confidence >= 0.0 && confidence <= 1.0, "Confidence should be in [0,1]"); + + Ok(()) + } + #[test] fn test_cosine_annealed_tau() { // Verify cosine schedule produces correct τ values at key points