From 1fdc712df606aea0f4729dee4b032a61aad83996 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Wed, 15 Apr 2026 07:01:19 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20trajectory=20backtracking=20=E2=80=94?= =?UTF-8?q?=20checkpoint=20rewind=20+=20informed=20perturbation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Detects Q-gap plateau (5 consecutive frozen epochs), rewinds to best-improvement checkpoint, applies perturbation strategy cycle: 1. Reset Adam momentum 2. Shrink-perturb branch heads 3. Temperature boost (2× for 5 epochs) 4. Learning rate boost (2× for 10 epochs) 3 rewinds × 4 routes = 12 max attempts. On exhaustion: save best model, exit with PLATEAU_EXHAUSTED. Saves ~3 hours H100 vs frozen training. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 52 ++++ crates/ml/src/trainers/dqn/fused_training.rs | 55 +++++ .../src/trainers/dqn/trainer/constructor.rs | 1 + crates/ml/src/trainers/dqn/trainer/mod.rs | 112 +++++++++ .../src/trainers/dqn/trainer/training_loop.rs | 226 ++++++++++++++++++ 5 files changed, 446 insertions(+) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 8dc152265..39f6a847c 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -3502,6 +3502,58 @@ impl GpuDqnTrainer { &self.target_params_buf } + /// Mutable reference to the flat f32 target parameter buffer. + pub fn target_params_mut(&mut self) -> &mut CudaSlice { + &mut self.target_params_buf + } + + /// Download online params from GPU to host (synchronous DtoH). + pub fn download_params(&self, dst: &mut [f32]) -> Result<(), MLError> { + unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } + self.stream.memcpy_dtoh(&self.params_buf, dst) + .map_err(|e| MLError::ModelError(format!("download params: {e}"))) + } + + /// Upload online params from host to GPU (HtoD). + pub fn upload_params(&mut self, src: &[f32]) -> Result<(), MLError> { + self.stream.memcpy_htod(src, &mut self.params_buf) + .map_err(|e| MLError::ModelError(format!("upload params: {e}"))) + } + + /// Download target params from GPU to host (synchronous DtoH). + pub fn download_target_params(&self, dst: &mut [f32]) -> Result<(), MLError> { + unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } + self.stream.memcpy_dtoh(&self.target_params_buf, dst) + .map_err(|e| MLError::ModelError(format!("download target params: {e}"))) + } + + /// Upload target params from host to GPU (HtoD). + pub fn upload_target_params(&mut self, src: &[f32]) -> Result<(), MLError> { + self.stream.memcpy_htod(src, &mut self.target_params_buf) + .map_err(|e| MLError::ModelError(format!("upload target params: {e}"))) + } + + /// Read eval_q_mean_ema. + pub fn eval_q_mean_ema(&self) -> f32 { self.eval_q_mean_ema } + + /// Read eval_q_std_ema. + pub fn eval_q_std_ema(&self) -> f32 { self.eval_q_std_ema } + + /// Read per_branch_q_gap_ema. + pub fn per_branch_q_gap_ema(&self) -> [f32; 4] { self.per_branch_q_gap_ema } + + /// Set eval_q_mean_ema (for trajectory backtracking restore). + pub fn set_eval_q_mean_ema(&mut self, v: f32) { self.eval_q_mean_ema = v; } + + /// Set eval_q_std_ema (for trajectory backtracking restore). + pub fn set_eval_q_std_ema(&mut self, v: f32) { self.eval_q_std_ema = v; } + + /// Set per_branch_q_gap_ema (for trajectory backtracking restore). + pub fn set_per_branch_q_gap_ema(&mut self, v: [f32; 4]) { self.per_branch_q_gap_ema = v; } + + /// Read adam_step counter. + pub fn adam_step_val(&self) -> i32 { self.adam_step } + // ═══════════════════════════════════════════════════════════════════ /// GPU-direct training step — no CPU roundtrip. diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 710e4876e..7e3d5356b 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -2110,6 +2110,61 @@ impl FusedTrainingCtx { self.trainer.total_params() } + // ── Trajectory backtracking accessors ────────────────────────────── + + /// Download online params from GPU to host (synchronous DtoH). + pub(crate) fn download_params(&self, dst: &mut [f32]) -> anyhow::Result<()> { + self.trainer.download_params(dst) + .map_err(|e| anyhow::anyhow!("download params: {e}")) + } + + /// Upload online params from host to GPU (HtoD). + pub(crate) fn upload_params(&mut self, src: &[f32]) -> anyhow::Result<()> { + self.trainer.upload_params(src) + .map_err(|e| anyhow::anyhow!("upload params: {e}")) + } + + /// Download target params from GPU to host (synchronous DtoH). + pub(crate) fn download_target_params(&self, dst: &mut [f32]) -> anyhow::Result<()> { + self.trainer.download_target_params(dst) + .map_err(|e| anyhow::anyhow!("download target params: {e}")) + } + + /// Upload target params from host to GPU (HtoD). + pub(crate) fn upload_target_params(&mut self, src: &[f32]) -> anyhow::Result<()> { + self.trainer.upload_target_params(src) + .map_err(|e| anyhow::anyhow!("upload target params: {e}")) + } + + /// Reset Adam optimizer state (m=0, v=0, step=0). + pub(crate) fn reset_adam_state(&mut self) -> anyhow::Result<()> { + self.trainer.reset_adam_state() + .map_err(|e| anyhow::anyhow!("reset adam: {e}")) + } + + /// Read adam_step counter. + pub(crate) fn adam_step_val(&self) -> i32 { self.trainer.adam_step_val() } + + /// Read eval_q_mean_ema. + pub(crate) fn eval_q_mean_ema(&self) -> f32 { self.trainer.eval_q_mean_ema() } + + /// Read eval_q_std_ema. + pub(crate) fn eval_q_std_ema(&self) -> f32 { self.trainer.eval_q_std_ema() } + + /// Read per_branch_q_gap_ema. + pub(crate) fn per_branch_q_gap_ema(&self) -> [f32; 4] { self.trainer.per_branch_q_gap_ema() } + + /// Set eval_q_mean_ema (trajectory backtracking restore). + pub(crate) fn set_eval_q_mean_ema(&mut self, v: f32) { self.trainer.set_eval_q_mean_ema(v); } + + /// Set eval_q_std_ema (trajectory backtracking restore). + pub(crate) fn set_eval_q_std_ema(&mut self, v: f32) { self.trainer.set_eval_q_std_ema(v); } + + /// Set per_branch_q_gap_ema (trajectory backtracking restore). + pub(crate) fn set_per_branch_q_gap_ema(&mut self, v: [f32; 4]) { self.trainer.set_per_branch_q_gap_ema(v); } + + // ── End trajectory backtracking accessors ────────────────────────── + /// Raw pointer + length of the current flat f32 params buffer. pub(crate) fn params_flat(&self) -> (u64, usize) { let p = self.trainer.params(); diff --git a/crates/ml/src/trainers/dqn/trainer/constructor.rs b/crates/ml/src/trainers/dqn/trainer/constructor.rs index ef5b84caf..63e4439a5 100644 --- a/crates/ml/src/trainers/dqn/trainer/constructor.rs +++ b/crates/ml/src/trainers/dqn/trainer/constructor.rs @@ -761,6 +761,7 @@ impl DQNTrainer { fold_val_start: 0, fold_val_end: 0, cached_n_episodes: None, + backtracking: super::BacktrackingState::new(), }) } diff --git a/crates/ml/src/trainers/dqn/trainer/mod.rs b/crates/ml/src/trainers/dqn/trainer/mod.rs index b4e5ede10..e690fae76 100644 --- a/crates/ml/src/trainers/dqn/trainer/mod.rs +++ b/crates/ml/src/trainers/dqn/trainer/mod.rs @@ -43,7 +43,116 @@ mod training_loop; #[cfg(test)] mod tests; +/// Saved training state for trajectory backtracking. +pub(crate) struct TrajectoryCheckpoint { + pub epoch: usize, + pub improvement_rate: f32, // d(val_Sharpe)/d(epoch) at save time + pub params_host: Vec, // DtoH copy of params_buf + pub target_params_host: Vec, // DtoH copy of target_params_buf + pub adam_step: i32, // optimizer step counter + pub eval_q_mean_ema: f32, + pub eval_q_std_ema: f32, + pub per_branch_q_gap_ema: [f32; 4], +} +/// Perturbation strategy for a rewind route. +#[derive(Debug, Clone, Copy)] +pub(crate) enum PerturbationStrategy { + /// Reset Adam momentum (m=0, v=0) -- fresh gradient estimates + ResetAdam, + /// Shrink-perturb branch heads only (sigma=0.1) + ShrinkPerturbBranches, + /// Temperature boost (2x Boltzmann + SARSA tau for 5 epochs, then cool) + TemperatureBoost, + /// Learning rate x2 for 10 epochs then restore + LearningRateBoost, +} + +/// Backtracking state machine. +pub(crate) struct BacktrackingState { + /// Top-3 checkpoints by improvement rate. + pub checkpoints: Vec, + /// Maximum checkpoints to keep. + pub max_checkpoints: usize, + /// Current rewind count (max 3). + pub rewind_count: usize, + /// Max rewinds before PLATEAU_EXHAUSTED. + pub max_rewinds: usize, + /// Current route index within a rewind (0-3). + pub route_idx: usize, + /// Routes per rewind (4 perturbation strategies). + pub routes_per_rewind: usize, + /// Consecutive plateau epochs (Q-gap velocity ~ 0). + pub plateau_epochs: usize, + /// Epochs required to trigger a rewind. + pub plateau_threshold: usize, + /// Minimum epochs between rewinds. + pub min_rewind_interval: usize, + /// Epoch of last rewind. + pub last_rewind_epoch: usize, + /// Whether a route is currently active (post-rewind trial). + pub route_active: bool, + /// Epoch when the current route started. + pub route_start_epoch: usize, + /// Epochs to evaluate a route before deciding. + pub route_eval_epochs: usize, + /// Val Sharpe at route start (to measure improvement). + pub route_start_sharpe: f32, + /// Minimum improvement rate to consider a route successful. + pub min_improvement_rate: f32, + /// Best val_Sharpe seen across all attempts. + pub best_sharpe: f32, + /// Epoch of best val_Sharpe. + pub best_epoch: usize, + /// Temperature boost active (decrements each epoch). + pub temp_boost_remaining: usize, + /// LR boost active (decrements each epoch). + pub lr_boost_remaining: usize, + /// Original LR (saved before LR boost). + pub original_lr: f64, +} + +impl BacktrackingState { + pub(crate) fn new() -> Self { + Self { + checkpoints: Vec::new(), + max_checkpoints: 3, + rewind_count: 0, + max_rewinds: 3, + route_idx: 0, + routes_per_rewind: 4, + plateau_epochs: 0, + plateau_threshold: 5, // 5 consecutive plateau epochs + min_rewind_interval: 10, + last_rewind_epoch: 0, + route_active: false, + route_start_epoch: 0, + route_eval_epochs: 10, + route_start_sharpe: 0.0, + min_improvement_rate: 0.1, // min 0.1 val_Sharpe/epoch improvement + best_sharpe: f32::NEG_INFINITY, + best_epoch: 0, + temp_boost_remaining: 0, + lr_boost_remaining: 0, + original_lr: 0.0, + } + } + + /// The 4 perturbation strategies, tried in order. + pub(crate) fn current_strategy(&self) -> PerturbationStrategy { + match self.route_idx % 4 { + 0 => PerturbationStrategy::ResetAdam, + 1 => PerturbationStrategy::ShrinkPerturbBranches, + 2 => PerturbationStrategy::TemperatureBoost, + _ => PerturbationStrategy::LearningRateBoost, + } + } + + /// Check if all rewinds and routes are exhausted. + pub(crate) fn is_exhausted(&self) -> bool { + self.rewind_count >= self.max_rewinds && self.route_idx >= self.routes_per_rewind + } +} pub struct DQNTrainer { /// DQN agent @@ -332,6 +441,9 @@ pub struct DQNTrainer { /// Cached GPU n_episodes (computed once, reused across epochs/folds). pub(crate) cached_n_episodes: Option, + + /// Task 10: Trajectory backtracking state machine for plateau recovery. + pub(crate) backtracking: BacktrackingState, } impl std::fmt::Debug for DQNTrainer { diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 5bcb4636e..ec1fa9064 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -539,6 +539,19 @@ impl DQNTrainer { &mut checkpoint_callback, ).await?; + // Task 10: Trajectory backtracking -- detect plateau, rewind, perturb + if self.run_backtracking_epoch_end(epoch, log_output.epoch_sharpe)? { + // PLATEAU_EXHAUSTED -- save best model and exit + info!( + "PLATEAU_EXHAUSTED: restoring best model from epoch {} before exit", + self.backtracking.best_epoch + ); + let _ = self.restore_best_gpu_params(); + let checkpoint_data = self.serialize_model().await?; + let _ = checkpoint_callback(epoch + 1, checkpoint_data, true); + break; + } + if (epoch + 1) % self.hyperparams.checkpoint_frequency == 0 { info!( "Saving periodic checkpoint at epoch {}/{}", @@ -2282,5 +2295,218 @@ impl DQNTrainer { Ok(()) } + + // ═══════════════════════════════════════════════════════════════════════ + // Task 10: Trajectory Backtracking helpers + // ═══════════════════════════════════════════════════════════════════════ + + /// Detect whether Q-gap is frozen (all branches have near-zero velocity). + fn is_q_gap_frozen(&self) -> bool { + if let Some(ref fused) = self.fused_ctx { + let emas = fused.per_branch_q_gap_ema(); + // Current q_gap approximated as uniform across branches (same as update_eval_v_range) + let q_gap = self.epoch_q_gap; + let current = [q_gap; 4]; + let threshold = 0.001_f32; + current.iter().zip(emas.iter()).all(|(gap, ema)| (gap - ema).abs() < threshold) + } else { + false + } + } + + /// Save a trajectory checkpoint when improvement is detected. + fn save_backtracking_checkpoint(&mut self, epoch: usize, improvement_rate: f32, val_sharpe: f32) -> anyhow::Result<()> { + if let Some(ref fused) = self.fused_ctx { + let total_params = fused.total_param_count(); + + // DtoH copy of params and target params + let mut params_host = vec![0.0_f32; total_params]; + let mut target_host = vec![0.0_f32; total_params]; + fused.download_params(&mut params_host)?; + fused.download_target_params(&mut target_host)?; + + let checkpoint = super::TrajectoryCheckpoint { + epoch, + improvement_rate, + params_host, + target_params_host: target_host, + adam_step: fused.adam_step_val(), + eval_q_mean_ema: fused.eval_q_mean_ema(), + eval_q_std_ema: fused.eval_q_std_ema(), + per_branch_q_gap_ema: fused.per_branch_q_gap_ema(), + }; + + // Insert sorted by improvement rate, keep top-3 + 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) + ); + if self.backtracking.checkpoints.len() > self.backtracking.max_checkpoints { + self.backtracking.checkpoints.pop(); + } + + info!(epoch, improvement_rate, checkpoints = self.backtracking.checkpoints.len(), + "Backtracking checkpoint saved (val_Sharpe={val_sharpe:.2})"); + } + Ok(()) + } + + /// Execute a rewind to the best checkpoint with the current perturbation strategy. + fn execute_rewind(&mut self, epoch: usize) -> anyhow::Result<()> { + let checkpoint_idx = self.backtracking.rewind_count.min(self.backtracking.checkpoints.len() - 1); + let strategy = self.backtracking.current_strategy(); + + // Log before borrowing checkpoint (strategy is Copy) + let ckpt_epoch = self.backtracking.checkpoints[checkpoint_idx].epoch; + info!( + epoch, + rewind_to = ckpt_epoch, + strategy = ?strategy, + rewind_count = self.backtracking.rewind_count, + route_idx = self.backtracking.route_idx, + "TRAJECTORY BACKTRACK: rewinding to epoch {} with {:?}", + ckpt_epoch, strategy + ); + + if let Some(ref mut fused) = self.fused_ctx { + // Upload params from checkpoint + fused.upload_params(&self.backtracking.checkpoints[checkpoint_idx].params_host)?; + fused.upload_target_params(&self.backtracking.checkpoints[checkpoint_idx].target_params_host)?; + + // Restore EMA state + let ema_mean = self.backtracking.checkpoints[checkpoint_idx].eval_q_mean_ema; + let ema_std = self.backtracking.checkpoints[checkpoint_idx].eval_q_std_ema; + let ema_gaps = self.backtracking.checkpoints[checkpoint_idx].per_branch_q_gap_ema; + fused.set_eval_q_mean_ema(ema_mean); + fused.set_eval_q_std_ema(ema_std); + fused.set_per_branch_q_gap_ema(ema_gaps); + + // Apply perturbation strategy + match strategy { + super::PerturbationStrategy::ResetAdam => { + fused.reset_adam_state()?; + info!("Perturbation: Adam momentum reset (m=0, v=0)"); + } + super::PerturbationStrategy::ShrinkPerturbBranches => { + fused.shrink_and_perturb(0.9, 0.1)?; + info!("Perturbation: shrink-perturb (alpha=0.9, sigma=0.1)"); + } + super::PerturbationStrategy::TemperatureBoost => { + self.backtracking.temp_boost_remaining = 5; + info!("Perturbation: temperature boost (2x for 5 epochs)"); + } + super::PerturbationStrategy::LearningRateBoost => { + self.backtracking.original_lr = self.hyperparams.learning_rate; + self.hyperparams.learning_rate *= 2.0; + self.backtracking.lr_boost_remaining = 10; + info!(lr = self.hyperparams.learning_rate, "Perturbation: LR boost (2x for 10 epochs)"); + } + } + } + + self.backtracking.last_rewind_epoch = epoch; + self.backtracking.route_active = true; + self.backtracking.route_start_epoch = epoch; + self.backtracking.route_start_sharpe = self.backtracking.best_sharpe; + self.backtracking.plateau_epochs = 0; + + Ok(()) + } + + /// Run backtracking logic at the end of each epoch. + /// Returns `true` if PLATEAU_EXHAUSTED — caller should break the epoch loop. + fn run_backtracking_epoch_end(&mut self, epoch: usize, val_sharpe: f64) -> anyhow::Result { + let val_sharpe_f32 = val_sharpe as f32; + + // 1. Track best Sharpe + if val_sharpe_f32 > self.backtracking.best_sharpe { + self.backtracking.best_sharpe = val_sharpe_f32; + self.backtracking.best_epoch = epoch; + } + + // 2. Save checkpoint when improving + let prev_sharpe = self.sharpe_history.last().copied().unwrap_or(0.0) as f32; + let improvement_rate = val_sharpe_f32 - prev_sharpe; // per epoch + if improvement_rate > 0.01 && !self.backtracking.route_active { + self.save_backtracking_checkpoint(epoch, improvement_rate, val_sharpe_f32)?; + } + + // 3. Detect plateau (Q-gap velocity ~ 0) + let q_gap_frozen = self.is_q_gap_frozen(); + if q_gap_frozen { + self.backtracking.plateau_epochs += 1; + } else { + self.backtracking.plateau_epochs = 0; + } + + // 4. Evaluate active route + if self.backtracking.route_active + && epoch.saturating_sub(self.backtracking.route_start_epoch) >= self.backtracking.route_eval_epochs + { + let route_improvement = (val_sharpe_f32 - self.backtracking.route_start_sharpe) + / self.backtracking.route_eval_epochs as f32; + + if route_improvement >= self.backtracking.min_improvement_rate { + // Route succeeded -- continue training normally + info!(epoch, route_improvement, "Backtracking route succeeded -- continuing"); + self.backtracking.route_active = false; + self.backtracking.plateau_epochs = 0; + } else { + // Route failed -- try next route or next rewind + info!(epoch, route_improvement, "Backtracking route failed -- trying next"); + self.backtracking.route_idx += 1; + + if self.backtracking.route_idx >= self.backtracking.routes_per_rewind { + // All routes for this rewind exhausted -- try next checkpoint + self.backtracking.rewind_count += 1; + self.backtracking.route_idx = 0; + } + + self.backtracking.route_active = false; + self.backtracking.plateau_epochs = self.backtracking.plateau_threshold; // re-trigger immediately + } + } + + // 5. Trigger rewind if plateau detected + if self.backtracking.plateau_epochs >= self.backtracking.plateau_threshold + && !self.backtracking.route_active + && epoch.saturating_sub(self.backtracking.last_rewind_epoch) >= self.backtracking.min_rewind_interval + && !self.backtracking.checkpoints.is_empty() + { + if self.backtracking.is_exhausted() { + // All attempts exhausted -- early termination + info!( + best_sharpe = self.backtracking.best_sharpe, + best_epoch = self.backtracking.best_epoch, + rewinds = self.backtracking.rewind_count, + "PLATEAU_EXHAUSTED: {}x{} = {} attempts. Best val_Sharpe={:.2} at epoch {}. Structural change needed.", + self.backtracking.max_rewinds, + self.backtracking.routes_per_rewind, + self.backtracking.max_rewinds * self.backtracking.routes_per_rewind, + self.backtracking.best_sharpe, + self.backtracking.best_epoch, + ); + return Ok(true); // signal caller to break + } + + // Rewind to best checkpoint + self.execute_rewind(epoch)?; + } + + // 6. Handle temperature/LR boost cooldown + if self.backtracking.temp_boost_remaining > 0 { + self.backtracking.temp_boost_remaining -= 1; + } + if self.backtracking.lr_boost_remaining > 0 { + self.backtracking.lr_boost_remaining -= 1; + if self.backtracking.lr_boost_remaining == 0 { + // Restore original LR + self.hyperparams.learning_rate = self.backtracking.original_lr; + info!(lr = self.hyperparams.learning_rate, "LR boost expired -- restored original learning rate"); + } + } + + Ok(false) + } }