diff --git a/crates/ml/src/dqn/distributional_dueling.rs b/crates/ml/src/dqn/distributional_dueling.rs index 09dd452fe..8f770f83d 100644 --- a/crates/ml/src/dqn/distributional_dueling.rs +++ b/crates/ml/src/dqn/distributional_dueling.rs @@ -44,6 +44,7 @@ use candle_nn::{Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; use crate::dqn::mixed_precision::training_dtype; +use crate::dqn::rmsnorm::RMSNorm; use crate::dqn::xavier_init::linear_xavier; use crate::MLError; @@ -122,14 +123,23 @@ pub struct DistributionalDuelingQNetwork { /// Shared feature extraction layers shared_layers: Vec, + /// RMSNorm after each shared hidden layer + shared_norms: Vec, + /// Value stream layers (outputs distribution) value_fc: Linear, value_out: Linear, // Output: [batch, num_atoms] + /// RMSNorm after value stream hidden layer + value_norm: RMSNorm, + /// Advantage stream layers (outputs distributions per action) advantage_fc: Linear, advantage_out: Linear, // Output: [batch, num_actions * num_atoms] + /// RMSNorm after advantage stream hidden layer + advantage_norm: RMSNorm, + /// Configuration config: DistributionalDuelingConfig, @@ -155,8 +165,9 @@ impl DistributionalDuelingQNetwork { let vars = VarMap::new(); let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); - // Build shared feature layers + // Build shared feature layers with RMSNorm after each let mut shared_layers = Vec::new(); + let mut shared_norms = Vec::new(); let mut current_dim = config.state_dim; for (i, &hidden_dim) in config.shared_hidden_dims.iter().enumerate() { @@ -166,6 +177,11 @@ impl DistributionalDuelingQNetwork { MLError::ModelError(format!("Failed to Xavier init shared layer {}: {}", i, e)) })?; shared_layers.push(layer); + + let norm_name = format!("shared_rmsnorm_{}", i); + let norm = RMSNorm::new_default(var_builder.pp(&norm_name), hidden_dim)?; + shared_norms.push(norm); + current_dim = hidden_dim; } @@ -173,6 +189,8 @@ impl DistributionalDuelingQNetwork { let value_fc_vb = var_builder.pp("value_fc"); let value_fc = linear_xavier(current_dim, config.value_hidden_dim, value_fc_vb) .map_err(|e| MLError::ModelError(format!("Failed to Xavier init value_fc: {}", e)))?; + let value_norm = + RMSNorm::new_default(var_builder.pp("value_rmsnorm"), config.value_hidden_dim)?; let value_out_vb = var_builder.pp("value_out"); let value_out = linear_xavier(config.value_hidden_dim, config.num_atoms, value_out_vb) @@ -184,6 +202,10 @@ impl DistributionalDuelingQNetwork { .map_err(|e| { MLError::ModelError(format!("Failed to Xavier init advantage_fc: {}", e)) })?; + let advantage_norm = RMSNorm::new_default( + var_builder.pp("advantage_rmsnorm"), + config.advantage_hidden_dim, + )?; let advantage_out_vb = var_builder.pp("advantage_out"); let advantage_out = linear_xavier( @@ -197,10 +219,13 @@ impl DistributionalDuelingQNetwork { Ok(Self { shared_layers, + shared_norms, value_fc, value_out, + value_norm, advantage_fc, advantage_out, + advantage_norm, config, vars, device, @@ -234,9 +259,14 @@ impl DistributionalDuelingQNetwork { .dim(0) .map_err(|e| MLError::ModelError(format!("Failed to get batch size: {}", e)))?; - // Shared feature extraction + // Shared feature extraction: Linear → LeakyReLU → RMSNorm let mut h = state.clone(); - for (i, layer) in self.shared_layers.iter().enumerate() { + for (i, (layer, norm)) in self + .shared_layers + .iter() + .zip(self.shared_norms.iter()) + .enumerate() + { h = layer.forward(&h).map_err(|e| { MLError::ModelError(format!("Shared layer {} forward failed: {}", i, e)) })?; @@ -245,26 +275,37 @@ impl DistributionalDuelingQNetwork { h = candle_nn::ops::leaky_relu(&h, self.config.leaky_relu_alpha).map_err(|e| { MLError::ModelError(format!("LeakyReLU failed at shared layer {}: {}", i, e)) })?; + + // RMSNorm stabilizes activations and gradients + h = norm.forward(&h).map_err(|e| { + MLError::ModelError(format!("RMSNorm failed at shared layer {}: {}", i, e)) + })?; } - // Value stream: V(s) → [batch, num_atoms] + // Value stream: Linear → LeakyReLU → RMSNorm → Linear → [batch, num_atoms] let v = self.value_fc.forward(&h).map_err(|e| { MLError::ModelError(format!("Value FC forward failed: {}", e)) })?; let v = candle_nn::ops::leaky_relu(&v, self.config.leaky_relu_alpha).map_err(|e| { MLError::ModelError(format!("Value LeakyReLU failed: {}", e)) })?; + let v = self.value_norm.forward(&v).map_err(|e| { + MLError::ModelError(format!("Value RMSNorm failed: {}", e)) + })?; let v = self.value_out.forward(&v).map_err(|e| { MLError::ModelError(format!("Value output forward failed: {}", e)) })?; // [batch, num_atoms] - // Advantage stream: A(s,a) → [batch, num_actions * num_atoms] + // Advantage stream: Linear → LeakyReLU → RMSNorm → Linear → [batch, num_actions * num_atoms] let a = self.advantage_fc.forward(&h).map_err(|e| { MLError::ModelError(format!("Advantage FC forward failed: {}", e)) })?; let a = candle_nn::ops::leaky_relu(&a, self.config.leaky_relu_alpha).map_err(|e| { MLError::ModelError(format!("Advantage LeakyReLU failed: {}", e)) })?; + let a = self.advantage_norm.forward(&a).map_err(|e| { + MLError::ModelError(format!("Advantage RMSNorm failed: {}", e)) + })?; let a_flat = self.advantage_out.forward(&a).map_err(|e| { MLError::ModelError(format!("Advantage output forward failed: {}", e)) })?; // [batch, num_actions * num_atoms]