fix: reset Adam optimizer state between walk-forward folds (prevents Q-value explosion)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-04 13:07:26 +02:00
parent 7354351a89
commit 3b6c9b5e4d
2 changed files with 30 additions and 0 deletions

View File

@@ -870,6 +870,33 @@ pub struct GpuDqnTrainer {
cql_d_adv_logits: CudaSlice<f32>,
}
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.

View File

@@ -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(())
}