diff --git a/crates/ml/src/trainers/dqn/trainer/mod.rs b/crates/ml/src/trainers/dqn/trainer/mod.rs index 5a9316599..1476e216d 100644 --- a/crates/ml/src/trainers/dqn/trainer/mod.rs +++ b/crates/ml/src/trainers/dqn/trainer/mod.rs @@ -46,6 +46,7 @@ mod tests; /// Saved training state for trajectory backtracking. pub(crate) struct TrajectoryCheckpoint { pub epoch: usize, + pub val_sharpe: f32, // val_Sharpe at save time — used for ranking pub improvement_rate: f32, // d(val_Sharpe)/d(epoch) at save time pub params_gpu: CudaSlice, // DtoD copy of params_buf — GPU resident pub target_params_gpu: CudaSlice, // DtoD copy of target_params_buf — GPU resident diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 3b69e8db8..0de3ab865 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -2358,6 +2358,7 @@ impl DQNTrainer { let checkpoint = super::TrajectoryCheckpoint { epoch, + val_sharpe, improvement_rate, params_gpu, target_params_gpu, @@ -2367,10 +2368,13 @@ impl DQNTrainer { per_branch_q_gap_ema: fused.per_branch_q_gap_ema(), }; - // Insert sorted by improvement rate, keep top-3 + // Sort by val_Sharpe (highest first), keep top-3. + // Rewind should go to the BEST state, not the fastest-improving. + // Epoch 0 (improvement=17.3, val_Sharpe=17) is useless as a rewind target. + // Epoch 20 (improvement=0.4, val_Sharpe=45) is the BEST rewind target. self.backtracking.checkpoints.push(checkpoint); self.backtracking.checkpoints.sort_by(|a, b| - b.improvement_rate.partial_cmp(&a.improvement_rate).unwrap_or(std::cmp::Ordering::Equal) + b.val_sharpe.partial_cmp(&a.val_sharpe).unwrap_or(std::cmp::Ordering::Equal) ); if self.backtracking.checkpoints.len() > self.backtracking.max_checkpoints { self.backtracking.checkpoints.pop();