From a16994b0ff8af945307936ad42db1616a673696a Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 7 Mar 2026 11:19:33 +0100 Subject: [PATCH] feat(ml): add RMSNorm to distributional-dueling network Add RMSNorm (root mean square normalization) after each hidden layer in the distributional-dueling Q-network: shared backbone layers, value stream, and advantage stream. RMSNorm stabilizes activations and gradients without the overhead of full LayerNorm (no mean centering), making the network less sensitive to input scale during training. Architecture per layer: Linear -> LeakyReLU -> RMSNorm RMSNorm weights are automatically tracked in the existing VarMap since they are created via VarBuilder::from_varmap with the same shared map. Co-Authored-By: Claude Opus 4.6 --- crates/ml/src/dqn/distributional_dueling.rs | 51 +++++++++++++++++++-- 1 file changed, 46 insertions(+), 5 deletions(-) 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]