diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 310cab788..2189ff22a 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -870,6 +870,33 @@ pub struct GpuDqnTrainer { cql_d_adv_logits: CudaSlice, } +impl GpuDqnTrainer { + /// Reset Adam optimizer state for a new walk-forward fold. + /// Zeroes momentum (m), variance (v), and step counter (t). + /// Weights are preserved (warm-start for the new fold). + pub fn reset_adam_state(&mut self) -> Result<(), MLError> { + self.stream.memset_zeros(&mut self.m_buf) + .map_err(|e| MLError::ModelError(format!("reset m_buf: {e}")))?; + self.stream.memset_zeros(&mut self.v_buf) + .map_err(|e| MLError::ModelError(format!("reset v_buf: {e}")))?; + self.stream.memset_zeros(&mut self.t_buf) + .map_err(|e| MLError::ModelError(format!("reset t_buf: {e}")))?; + self.adam_step = 0; + + // Also reset IQN trunk Adam state (currently unused but allocated). + self.stream.memset_zeros(&mut self.iqn_trunk_m) + .map_err(|e| MLError::ModelError(format!("reset iqn m: {e}")))?; + self.stream.memset_zeros(&mut self.iqn_trunk_v) + .map_err(|e| MLError::ModelError(format!("reset iqn v: {e}")))?; + self.stream.memset_zeros(&mut self.iqn_trunk_t_buf) + .map_err(|e| MLError::ModelError(format!("reset iqn t: {e}")))?; + self.iqn_trunk_adam_step = 0; + + tracing::info!("Adam optimizer state reset for new fold"); + Ok(()) + } +} + impl Drop for GpuDqnTrainer { fn drop(&mut self) { // Synchronize stream and destroy graphs BEFORE CudaSlice fields drop. diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index f1f42e3db..20650ebea 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -584,6 +584,9 @@ impl FusedTrainingCtx { self.steps_since_varmap_sync = 0; self.last_combined_norm = 0.0; self.graph_aux = None; + // Reset Adam optimizer state — each fold starts with fresh momentum. + // Stale momentum from a previous fold causes weight explosion → Q-value -1e30. + self.trainer.reset_adam_state()?; Ok(()) }