diff --git a/ml/src/ppo/lstm_networks.rs b/ml/src/ppo/lstm_networks.rs index 2722f19c6..316abd6cd 100644 --- a/ml/src/ppo/lstm_networks.rs +++ b/ml/src/ppo/lstm_networks.rs @@ -162,6 +162,61 @@ impl LSTMPolicyNetwork { Ok((logits, new_h, new_c)) } + /// Reconstruct an LSTM policy network from a VarBuilder (checkpoint loading) + /// + /// This loads network weights from a safetensors checkpoint. LSTM hidden/cell + /// states are NOT saved in checkpoints -- they are re-initialized to zeros on load. + /// + /// # Arguments + /// * `config` - PPO configuration (must match the config used during training) + /// * `vb` - VarBuilder backed by loaded safetensors checkpoint + /// * `device` - Device to load model on + /// + /// # Returns + /// LSTMPolicyNetwork with restored weights, or error if checkpoint structure + /// doesn't match config + pub fn from_varbuilder( + config: &super::ppo::PPOConfig, + vb: VarBuilder<'_>, + device: &Device, + ) -> Result { + let hidden_dim = config.lstm_hidden_dim; + let num_layers = config.lstm_num_layers; + let input_dim = config.state_dim; + let output_dim = config.num_actions; + + // Load input projection layer (state_dim -> hidden_dim) + let input_layer = linear(input_dim, hidden_dim, vb.pp("input"))?; + + // Load LSTM layers + let mut lstm_layers = Vec::new(); + for layer_idx in 0..num_layers { + let lstm_config = LSTMConfig { + layer_idx, + ..Default::default() + }; + let in_dim = if layer_idx == 0 { hidden_dim } else { hidden_dim }; + let lstm = LSTM::new(in_dim, hidden_dim, lstm_config, vb.pp(format!("lstm_{}", layer_idx)))?; + lstm_layers.push(lstm); + } + + // Load output layer (hidden_dim -> num_actions) + let output_layer = linear(hidden_dim, output_dim, vb.pp("output"))?; + + // Create empty VarMap (weights are in VarBuilder, not VarMap for loaded models) + let vars = VarMap::new(); + + Ok(Self { + input_layer, + lstm_layers, + output_layer, + device: device.clone(), + vars, + hidden_dim, + num_layers, + }) + } + /// Get network variables pub fn vars(&self) -> &VarMap { &self.vars @@ -326,6 +381,60 @@ impl LSTMValueNetwork { Ok((value, new_h, new_c)) } + /// Reconstruct an LSTM value network from a VarBuilder (checkpoint loading) + /// + /// This loads network weights from a safetensors checkpoint. LSTM hidden/cell + /// states are NOT saved in checkpoints -- they are re-initialized to zeros on load. + /// + /// # Arguments + /// * `config` - PPO configuration (must match the config used during training) + /// * `vb` - VarBuilder backed by loaded safetensors checkpoint + /// * `device` - Device to load model on + /// + /// # Returns + /// LSTMValueNetwork with restored weights, or error if checkpoint structure + /// doesn't match config + pub fn from_varbuilder( + config: &super::ppo::PPOConfig, + vb: VarBuilder<'_>, + device: &Device, + ) -> Result { + let hidden_dim = config.lstm_hidden_dim; + let num_layers = config.lstm_num_layers; + let input_dim = config.state_dim; + + // Load input projection layer (state_dim -> hidden_dim) + let input_layer = linear(input_dim, hidden_dim, vb.pp("input"))?; + + // Load LSTM layers + let mut lstm_layers = Vec::new(); + for layer_idx in 0..num_layers { + let lstm_config = LSTMConfig { + layer_idx, + ..Default::default() + }; + let in_dim = if layer_idx == 0 { hidden_dim } else { hidden_dim }; + let lstm = LSTM::new(in_dim, hidden_dim, lstm_config, vb.pp(format!("lstm_{}", layer_idx)))?; + lstm_layers.push(lstm); + } + + // Load output layer (hidden_dim -> 1) + let output_layer = linear(hidden_dim, 1, vb.pp("output"))?; + + // Create empty VarMap (weights are in VarBuilder, not VarMap for loaded models) + let vars = VarMap::new(); + + Ok(Self { + input_layer, + lstm_layers, + output_layer, + device: device.clone(), + vars, + hidden_dim, + num_layers, + }) + } + /// Get network variables pub fn vars(&self) -> &VarMap { &self.vars diff --git a/ml/src/ppo/ppo.rs b/ml/src/ppo/ppo.rs index 7eb4dbd7e..351fc3adf 100644 --- a/ml/src/ppo/ppo.rs +++ b/ml/src/ppo/ppo.rs @@ -18,7 +18,7 @@ use candle_optimisers::adam::ParamsAdam; use rand::{thread_rng, Rng}; use serde::{Deserialize, Serialize}; use std::path::PathBuf; -use tracing::{info, warn}; +use tracing::{debug, info, warn}; use crate::tensor_ops::TensorOps; @@ -927,22 +927,54 @@ impl WorkingPPO { } } - // Update policy network - // Note: Gradient clipping is not available in candle 0.9.1 API - // Instead, we rely on reduced learning rate (3e-5) to prevent gradient explosion + // Update policy network with gradient norm monitoring + let actor_vars = self.actor.vars().all_vars(); + let policy_grads = policy_loss.backward().map_err(|e| { + MLError::TrainingError(format!("Policy backward failed: {}", e)) + })?; + + let policy_grad_norm = + Self::compute_gradient_norm(&actor_vars, &policy_grads)?; + if let Some(ref mut optimizer) = self.policy_optimizer { - optimizer.backward_step(&policy_loss).map_err(|e| { - MLError::TrainingError(format!("Policy backward step failed: {}", e)) + optimizer.step(&policy_grads).map_err(|e| { + MLError::TrainingError(format!("Policy optimizer step failed: {}", e)) })?; } - // Update value network + if policy_grad_norm > self.config.max_grad_norm as f64 { + warn!( + "Policy gradient norm {:.4} exceeds max {:.4}. Consider reducing learning rate.", + policy_grad_norm, self.config.max_grad_norm + ); + } + + debug!("Policy gradient norm: {:.4}", policy_grad_norm); + + // Update value network with gradient norm monitoring + let critic_vars = self.critic.vars().all_vars(); + let value_grads = value_loss.backward().map_err(|e| { + MLError::TrainingError(format!("Value backward failed: {}", e)) + })?; + + let value_grad_norm = + Self::compute_gradient_norm(&critic_vars, &value_grads)?; + if let Some(ref mut optimizer) = self.value_optimizer { - optimizer.backward_step(&value_loss).map_err(|e| { - MLError::TrainingError(format!("Value backward step failed: {}", e)) + optimizer.step(&value_grads).map_err(|e| { + MLError::TrainingError(format!("Value optimizer step failed: {}", e)) })?; } + if value_grad_norm > self.config.max_grad_norm as f64 { + warn!( + "Value gradient norm {:.4} exceeds max {:.4}. Consider reducing learning rate.", + value_grad_norm, self.config.max_grad_norm + ); + } + + debug!("Value gradient norm: {:.4}", value_grad_norm); + total_policy_loss += policy_loss_scalar; total_value_loss += value_loss_scalar; num_updates += 1; @@ -1163,19 +1195,54 @@ impl WorkingPPO { } } - // Update networks + // Update policy network with gradient norm monitoring + let actor_vars = self.actor.vars().all_vars(); + let policy_grads = policy_loss.backward().map_err(|e| { + MLError::TrainingError(format!("Policy backward failed: {}", e)) + })?; + + let policy_grad_norm = + Self::compute_gradient_norm(&actor_vars, &policy_grads)?; + if let Some(ref mut optimizer) = self.policy_optimizer { - optimizer.backward_step(&policy_loss).map_err(|e| { - MLError::TrainingError(format!("Policy backward step failed: {}", e)) + optimizer.step(&policy_grads).map_err(|e| { + MLError::TrainingError(format!("Policy optimizer step failed: {}", e)) })?; } + if policy_grad_norm > self.config.max_grad_norm as f64 { + warn!( + "LSTM policy gradient norm {:.4} exceeds max {:.4}. Consider reducing learning rate or sequence length.", + policy_grad_norm, self.config.max_grad_norm + ); + } + + debug!("LSTM policy gradient norm: {:.4}", policy_grad_norm); + + // Update value network with gradient norm monitoring + let critic_vars = self.critic.vars().all_vars(); + let value_grads = scaled_value_loss.backward().map_err(|e| { + MLError::TrainingError(format!("Value backward failed: {}", e)) + })?; + + let value_grad_norm = + Self::compute_gradient_norm(&critic_vars, &value_grads)?; + if let Some(ref mut optimizer) = self.value_optimizer { - optimizer.backward_step(&scaled_value_loss).map_err(|e| { - MLError::TrainingError(format!("Value backward step failed: {}", e)) + optimizer.step(&value_grads).map_err(|e| { + MLError::TrainingError(format!("Value optimizer step failed: {}", e)) })?; } + if value_grad_norm > self.config.max_grad_norm as f64 { + warn!( + "LSTM value gradient norm {:.4} exceeds max {:.4}. Consider reducing learning rate.", + value_grad_norm, self.config.max_grad_norm + ); + } + + debug!("LSTM value gradient norm: {:.4}", value_grad_norm); + total_policy_loss += policy_loss_scalar; total_value_loss += value_loss_scalar; num_updates += 1; @@ -1200,6 +1267,43 @@ impl WorkingPPO { Ok((avg_policy_loss, avg_value_loss)) } + /// Compute the L2 norm of all gradients for monitoring + /// + /// This follows the same pattern as `ContinuousPPO::compute_gradient_norm`. + /// The norm is used to detect gradient explosion; when it exceeds + /// `max_grad_norm`, a warning is logged recommending learning rate reduction. + fn compute_gradient_norm( + vars: &[candle_core::Var], + grads: &candle_core::backprop::GradStore, + ) -> Result { + let mut total_norm_sq = 0.0f64; + + for var in vars { + if let Some(grad) = grads.get(var) { + let grad_norm_sq = grad + .sqr() + .map_err(|e| { + MLError::TrainingError(format!("Failed to square gradient: {}", e)) + })? + .sum_all() + .map_err(|e| { + MLError::TrainingError(format!("Failed to sum gradient: {}", e)) + })? + .to_vec0::() + .map_err(|e| { + MLError::TrainingError(format!( + "Failed to extract gradient norm: {}", + e + )) + })? as f64; + + total_norm_sq += grad_norm_sq; + } + } + + Ok(total_norm_sq.sqrt()) + } + /// Compute losses WITHOUT updating network weights (for validation) /// /// This method is used during hyperparameter optimization to compute @@ -1438,16 +1542,6 @@ impl WorkingPPO { let lstm_hidden_dim = config.lstm_hidden_dim; let lstm_num_layers = config.lstm_num_layers; - // TODO: Implement from_varbuilder for LSTM networks to enable checkpoint loading - // For now, only MLP checkpoints are supported - if use_lstm { - return Err(MLError::ModelError( - "Loading LSTM checkpoints is not yet implemented. \ - LSTM networks require from_varbuilder methods in lstm_networks.rs. \ - Use with_device() constructor for LSTM mode instead.".to_string() - )); - } - // Load actor network from safetensors let actor_path = PathBuf::from(actor_checkpoint_path); // SAFETY: VarBuilder::from_mmaped_safetensors is safe here because: @@ -1482,15 +1576,28 @@ impl WorkingPPO { )? }; - // Load MLP actor network (LSTM checkpoint loading not yet supported) - let mlp_actor = PolicyNetwork::from_varbuilder( - actor_vb, - config.state_dim, - &config.policy_hidden_dims, - config.num_actions, - device.clone(), - )?; - let actor = ActorNetwork::MLP(mlp_actor); + let actor = if use_lstm { + info!( + "Loading LSTM actor checkpoint (hidden_dim={}, num_layers={})", + lstm_hidden_dim, lstm_num_layers + ); + let lstm_actor = LSTMPolicyNetwork::from_varbuilder(&config, actor_vb, &device) + .map_err(|e| { + MLError::ModelError(format!( + "Failed to load LSTM actor from checkpoint: {}", e + )) + })?; + ActorNetwork::LSTM(lstm_actor) + } else { + let mlp_actor = PolicyNetwork::from_varbuilder( + actor_vb, + config.state_dim, + &config.policy_hidden_dims, + config.num_actions, + device.clone(), + )?; + ActorNetwork::MLP(mlp_actor) + }; // Load critic network from safetensors let critic_path = PathBuf::from(critic_checkpoint_path); @@ -1526,14 +1633,27 @@ impl WorkingPPO { )? }; - // Load MLP critic network (LSTM checkpoint loading not yet supported) - let mlp_critic = ValueNetwork::from_varbuilder( - critic_vb, - config.state_dim, - &config.value_hidden_dims, - device.clone(), - )?; - let critic = CriticNetwork::MLP(mlp_critic); + let critic = if use_lstm { + info!( + "Loading LSTM critic checkpoint (hidden_dim={}, num_layers={})", + lstm_hidden_dim, lstm_num_layers + ); + let lstm_critic = LSTMValueNetwork::from_varbuilder(&config, critic_vb, &device) + .map_err(|e| { + MLError::ModelError(format!( + "Failed to load LSTM critic from checkpoint: {}", e + )) + })?; + CriticNetwork::LSTM(lstm_critic) + } else { + let mlp_critic = ValueNetwork::from_varbuilder( + critic_vb, + config.state_dim, + &config.value_hidden_dims, + device.clone(), + )?; + CriticNetwork::MLP(mlp_critic) + }; // Try to load metadata to restore training_steps // Look for metadata file in same directory as actor checkpoint @@ -1611,7 +1731,8 @@ impl WorkingPPO { // Note: use_lstm, lstm_hidden_dim, lstm_num_layers already extracted at function start // Initialize hidden state manager if LSTM is enabled - // (LSTM checkpoints not supported yet, so this will always be None for loaded checkpoints) + // LSTM hidden/cell states are not saved in checkpoints -- only network weights. + // States are re-initialized to zeros on load via HiddenStateManager::new(). let hidden_state_manager = if use_lstm { Some(HiddenStateManager::new( lstm_num_layers,