From a01a376bd2e853911ae2a1e12b0ef01cd777086f Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 24 May 2026 10:59:50 +0200 Subject: [PATCH] audit: split K-loop to DQN-only (avoid PPO/V overshoot at high K) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per `pearl_q_thompson_actor_makes_pi_dead_weight` follow-up + #35 deferral: the K-loop in step_with_lobsim was running full step_synthetic K times per env step (Q + π + V + encoder + LR controller emit + ISV refresh + EMA inputs). At K=4 (default) or K=8 (prior K_MAX) this caused PPO overshoot — KL excursions to 12.44 in f2ggr, policy drift faster than the env step rate, gradient overtraining on the same env-step's h_t. ## Fix: extract dqn_replay_step helper New public method `dqn_replay_step(b_size)`: 1. Forward Q on sampled_h_t + sampled_h_tp1 (Double-DQN argmax) 2. Bellman target via TARGET net at h_tp1 + select + project 3. Q backward (logits → grad_w/b/h_t) 4. Per-batch reduce → grad_w/grad_b 5. Q Adam (uses LR already set by step_synthetic — no re-fire of the LR controller per K iter) 6. Writes td_per_sample_d for PER priority update by caller Discards Q's grad_h_t per R7d stop-grad (same as step_synthetic). What dqn_replay_step does NOT do: * π forward / surrogate / Adam — runs once per env step in step_synthetic * V forward / backward / Adam — same * Encoder backward / grad combine — same * LR controller emit + ISV mirror refresh — same * EMA inputs (entropy, KL, advantage_var, td_kurtosis) — same ## K-loop in step_with_lobsim for k_iter in 0..k_updates { let per_indices = sample_and_gather(b_size)?; if k_iter == 0 { stats = step_synthetic(snapshots)?; // full update } else { dqn_replay_step(b_size)?; // Q-only } // PER priority update } Result: * Q gets K Adam updates per env step (K-fold variance reduction) * π + V + encoder get 1 Adam update per env step (no overshoot) * LR controllers fire once per env step (no double-counting of plateau detection) * At b_size=16 with low advantage_var_ratio (batch averaging reduces noise), K-loop typically settles at K=1 — the split becomes a no-op in the steady state. At b_size=1 fallback or high-noise regimes, the split materially reduces PPO drift. ## Code duplication dqn_replay_step duplicates ~120 lines of Q-section code from step_synthetic. Acceptable temporary tech debt — full dedupe would require restructuring step_synthetic to call dqn_replay_step internally, which is a larger refactor with regression risk. Marked TODO for a follow-up commit once the b_size=16 + π-actor + K-split architecture is empirically validated. ## Verified gates (local sm_86) G1 isv_bootstrap ✅ G3 controllers ✅ G4 target_update ✅ integrated_smoke ✅ ## No smoke yet alpha-rl-9k9x6 (commit 3737feb66, π-actor + b_size=16) is still in flight; submitting a new smoke would compete for L40S GPU. This commit lands on remote; smoke will be submitted after 9k9x6 lands and we've analyzed whether the b_size=16 + π-actor architecture worked. If 9k9x6 shows Q learning unblocked, this split is polish. If it doesn't, the split becomes the next experiment. Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/src/trainer/integrated.rs | 193 ++++++++++++++++++++-- 1 file changed, 183 insertions(+), 10 deletions(-) diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index 512b69195..efcc88b89 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -2322,6 +2322,166 @@ impl IntegratedTrainer { }) } + /// DQN-only replay step — used by step_with_lobsim's K-loop for + /// iterations after the first. Skips PPO surrogate, V regression, + /// encoder backward, LR controller emit, ISV mirror refresh — all + /// of those run once per env step in the first iter via + /// step_synthetic. This helper just does: + /// + /// 1. Forward Q on `sampled_h_t` and `sampled_h_tp1` + /// 2. Double-DQN argmax on online Q at h_tp1 → + /// `sampled_next_actions_d` + /// 3. Bellman target via `forward_target` (target net) + + /// `select_action_atoms` + `project_bellman_target` + /// 4. Q backward: logits → grad_w/b/h_t. Q's grad_h_t is + /// discarded (R7d stop-grad: SAMPLED h_t is a past-step + /// encoder output; folding its gradient into the encoder + /// would poison live training). + /// 5. Per-batch grad reduce → grad_w/grad_b + /// 6. Q Adam (uses LR already set in step_synthetic via LR + /// controller emit) + /// 7. Writes `td_per_sample_d` for PER priority update by the + /// caller. + /// + /// Assumes `sample_and_gather` has been called externally — + /// reads `sampled_h_t_d / sampled_h_tp1_d / sampled_actions_d / + /// sampled_rewards_d / sampled_dones_d` and writes + /// `sampled_next_actions_d`. + /// + /// Returns mean Q loss (scalar). Caller uses this for the + /// LR-controller's plateau detection (but DOES NOT mirror it into + /// `last_q_loss` — that field is the env-step-paired iter's Q + /// loss, the canonical reference for plateau detection). + pub fn dqn_replay_step(&mut self, b_size: usize) -> Result { + if b_size == 0 { + anyhow::bail!("dqn_replay_step: empty batch (b_size = 0)"); + } + debug_assert_eq!(self.sampled_h_t_d.len(), b_size * HIDDEN_DIM); + debug_assert_eq!(self.sampled_h_tp1_d.len(), b_size * HIDDEN_DIM); + + let k_dqn = N_ACTIONS * Q_N_ATOMS; + + // Per-iter scratch — small relative to the GPU allocator + // amortisation, and freed when the function returns. + let mut q_logits_d = self.stream.alloc_zeros::(b_size * k_dqn)?; + let mut q_logits_tp1_sampled_d = self.stream.alloc_zeros::(b_size * k_dqn)?; + let q_loss_d = self.stream.alloc_zeros::(1)?; + let mut q_target_full_d = self.stream.alloc_zeros::(b_size * k_dqn)?; + let mut q_target_action_d = self.stream.alloc_zeros::(b_size * Q_N_ATOMS)?; + let mut target_dist_d = self.stream.alloc_zeros::(b_size * Q_N_ATOMS)?; + let mut q_grad_logits_d = self.stream.alloc_zeros::(b_size * k_dqn)?; + let mut q_grad_w_per_batch_d = self + .stream + .alloc_zeros::(b_size * k_dqn * HIDDEN_DIM)?; + let mut q_grad_b_per_batch_d = self.stream.alloc_zeros::(b_size * k_dqn)?; + let mut q_grad_h_t_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; + let mut q_grad_w_d = self.stream.alloc_zeros::(k_dqn * HIDDEN_DIM)?; + let mut q_grad_b_d = self.stream.alloc_zeros::(k_dqn)?; + + let b_size_i = b_size as i32; + + // ── 1. Forward Q on sampled_h_t and sampled_h_tp1 ─────────── + self.dqn_head + .forward(&self.sampled_h_t_d, b_size, &mut q_logits_d) + .context("dqn_replay_step: dqn_head.forward(sampled_h_t)")?; + self.dqn_head + .forward(&self.sampled_h_tp1_d, b_size, &mut q_logits_tp1_sampled_d) + .context("dqn_replay_step: dqn_head.forward(sampled_h_tp1)")?; + + // ── 2. Double-DQN argmax on online Q at h_tp1 ─────────────── + { + let cfg_argmax = LaunchConfig { + grid_dim: (b_size as u32, 1, 1), + block_dim: (N_ACTIONS as u32, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.argmax_expected_q_fn); + launch + .arg(&q_logits_tp1_sampled_d) + .arg(&self.atom_supports_d) + .arg(&mut self.sampled_next_actions_d) + .arg(&b_size_i); + unsafe { + launch + .launch(cfg_argmax) + .context("dqn_replay_step: argmax_expected_q launch")?; + } + } + + // ── 3. Bellman target build via TARGET net at h_tp1 ───────── + self.dqn_head + .forward_target(&self.sampled_h_tp1_d, b_size, &mut q_target_full_d) + .context("dqn_replay_step: dqn_head.forward_target(sampled_h_tp1)")?; + self.dqn_head + .select_action_atoms( + &q_target_full_d, + &self.sampled_next_actions_d, + b_size, + &mut q_target_action_d, + ) + .context("dqn_replay_step: dqn_head.select_action_atoms")?; + self.dqn_head + .project_bellman_target( + &q_target_action_d, + &self.sampled_rewards_d, + &self.sampled_dones_d, + &self.isv_d, + b_size, + &mut target_dist_d, + ) + .context("dqn_replay_step: dqn_head.project_bellman_target")?; + + // ── 4. Q backward (logits → grad_w/b/h_t) ─────────────────── + let mut q_loss_d_mut = q_loss_d; + self.dqn_head + .backward_logits( + &q_logits_d, + &target_dist_d, + &self.sampled_actions_d, + b_size, + &mut q_loss_d_mut, + &mut self.td_per_sample_d, + &mut q_grad_logits_d, + ) + .context("dqn_replay_step: dqn_head.backward_logits")?; + let l_q_host = read_scalar_d(&self.stream, &q_loss_d_mut)?; + let l_q = l_q_host / (b_size as f32); + + self.dqn_head + .backward_to_w_b_h( + &self.sampled_h_t_d, + &q_grad_logits_d, + b_size, + &mut q_grad_w_per_batch_d, + &mut q_grad_b_per_batch_d, + &mut q_grad_h_t_d, + ) + .context("dqn_replay_step: dqn_head.backward_to_w_b_h")?; + // R7d stop-grad: discard q_grad_h_t (sampled h_t is past-step + // encoder output; folding its gradient into the encoder would + // poison live training). Same semantics as step_synthetic. + let _ = &q_grad_h_t_d; + + self.launch_reduce_axis0( + &q_grad_w_per_batch_d, + b_size, + k_dqn * HIDDEN_DIM, + &mut q_grad_w_d, + )?; + self.launch_reduce_axis0(&q_grad_b_per_batch_d, b_size, k_dqn, &mut q_grad_b_d)?; + + // ── 5. Q Adam (uses LR set by step_synthetic; we don't re-fire + // the LR controller here — that runs once per env step). + self.dqn_w_adam + .step(&mut self.dqn_head.w_d, &q_grad_w_d) + .context("dqn_replay_step: dqn_w_adam.step")?; + self.dqn_b_adam + .step(&mut self.dqn_head.b_d, &q_grad_b_d) + .context("dqn_replay_step: dqn_b_adam.step")?; + + Ok(l_q) + } + /// Phase R6: run one integrated training step driven by a real /// `LobSimCuda` (via the narrow `RlLobBackend` trait, dep-cycle /// break only). GPU-pure for the env interaction: Thompson- @@ -2935,21 +3095,34 @@ impl IntegratedTrainer { .clamp(1, k_max); self.last_k_updates = k_updates; - // Track the LAST iteration's stats for caller reporting (per- - // iteration stats are not aggregated — caller sees the loss - // after the final Adam step of the last iteration, which is - // the most-recent state of the model). + // K-loop split (Path B): first iter does the full env-step- + // paired Q+π+V+encoder update via step_synthetic. Subsequent + // iters do DQN-ONLY replay via dqn_replay_step — fresh PER + // sample + Q forward + Bellman + Q backward + Q Adam. π+V+ + // encoder updates stay at 1× per env step regardless of K, so + // the K-loop intensifies Q learning without overshooting PPO + // (avoids the f2ggr KL=12.44 pathology where K=8 ran 8× PPO + // updates per env-step). At b_size=16 the K-loop typically + // settles at K=1 (advantage_var_ratio drops with batch size), + // but the split protects against pathological regimes and + // makes the actor/critic separation explicit. let mut stats = IntegratedStepStats::default(); - for _k_iter in 0..k_updates { + for k_iter in 0..k_updates { // Fresh PER sample per iteration — different transitions - // each gradient step. Without this the K-loop would just - // do K Adam steps on the same gradient, equivalent to a - // single step with K × LR (and would interact badly with - // Adam's per-parameter moment estimates). + // each gradient step. let per_indices = self .sample_and_gather(b_size) .context("step_with_lobsim: sample_and_gather (k-loop)")?; - stats = self.step_synthetic(snapshots)?; + if k_iter == 0 { + // First iter: full step_synthetic (env-step-paired). + stats = self.step_synthetic(snapshots)?; + } else { + // Subsequent iters: Q-only replay (no π, no V, no + // encoder backward, no LR controller refresh). + let _l_q_extra = self + .dqn_replay_step(b_size) + .context("step_with_lobsim: dqn_replay_step (k-loop iter > 0)")?; + } // PER priority update per iteration so the next sample // picks transitions with updated TD errors. if !per_indices.is_empty() {