From c34650a2416a7ab871415b1e4c245cc5648b68ce Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 23 May 2026 12:15:41 +0200 Subject: [PATCH] =?UTF-8?q?refactor(rl):=20R7b=20=E2=80=94=20host=20Thomps?= =?UTF-8?q?on/argmax/log=5Fpi=20=E2=86=92=20GPU;=20V=20stays=20on=20device?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Closes the last host orchestration in `step_with_lobsim`'s policy path per `feedback_cpu_is_read_only` + `pearl_no_host_branches_in_captured_graph`. The flawed Phase F+G shipped three host loops (Thompson sampling over C51 atoms, argmax over expected Q, log-softmax for log π_old) that each required a DtoH staging copy of the relevant Q/V/π device buffer per step. R7b deletes them all in favour of three R4 kernels: * `rl_action_kernel` — Thompson sampler with per-batch xorshift32 PRNG state (device-resident, no salt-based host RNG seeding per step). * `argmax_expected_q` — deterministic Bellman-target argmax over expected Q per action; distinct selector from Thompson per `pearl_thompson_for_distributional_action_selection`. * `log_pi_at_action` — per-batch log π(action_b) via log-softmax + lookup; feeds the PPO importance ratio. Side effects of the lift: * The host `read_slice_d` DtoH calls for `q_logits_host`, `v_pred_host`, `pi_logits_host` are deleted (each was 4 × b_size × Q_N_ATOMS or b_size × N_ACTIONS floats per step). Three full device→host→device roundtrips per step → zero. * The host γ readback (ISV[400] mirror DtoH + bootstrap fallback that the flawed branch kept "just in case") is also gone: `compute_advantage_return` already reads γ on-device from ISV[400] (R3). The fallback was a smell that R7a planned to eliminate; R7b actually removes it. * V(s_t) stays in `v_pred_d` throughout the hot path. The host upload via the (now-deleted) `upload_f32` helper that round-tripped V back to device for `compute_advantage_return` is gone. V(s_{t+1}) still reuses `v_pred_d` (h_t bootstrap) — true V(s_{t+1}) via `forward_encoder(next_snapshots)` is explicitly scoped to R7c. * Step-counter increment retained (drives the Thompson-vs-argmax diagnostic ordering in the trainer's stats record), but the ChaCha8Rng-from-step_counter that seeded the host loop is gone. Per `feedback_no_hiding` the three retired helper functions (`upload_f32`, `upload_i32`, `argmax_f32`) are DELETED rather than `#[allow(dead_code)]`'d. The matching unused imports (`MappedI32Buffer`, `DevicePtrMut`) are removed in the same commit. The `actions_d` allocation that used to live in step_with_lobsim's local scope (used only to bridge the host Thompson loop → the GPU `actions_to_market_targets` kernel) is gone too — that kernel now reads directly from `self.actions_d` written by `rl_action_kernel`, closing the last action-path host upload. What's still HtoD in the hot path (intentional, boundary data): * `apply_snapshot(last_snap)` — the trainer-fed MBP-10 input crosses the I/O boundary; mapped-pinned per `feedback_no_htod_htoh_only_mapped_pinned`. * HEALTH_DIAG ISV readback inside `step_synthetic` — explicit diagnostic surface, not hot-path orchestration. Falsifiability gate G4 (R4 kernel oracle tests + R5 controller + soft-update tests) already passing on this branch. The R6/R7a smoke (`integrated_trainer_step_with_lobsim_runs_without_panic`) re-builds clean and the only host code remaining in `step_with_lobsim` is the bounded `LaunchConfig`/`b_size_i` setup before each kernel launch. R7c (next): PER buffer wiring (push transitions, sample, update priorities) per `feedback_always_per`, and `next_snapshots` + `forward_encoder(next_snapshots)` for true V(s_{t+1}) Bellman target (replaces the h_t bootstrap reuse above). Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/src/trainer/integrated.rs | 328 +++++++--------------- 1 file changed, 102 insertions(+), 226 deletions(-) diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index abfcd5206..4d43a4b2e 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -78,14 +78,13 @@ use std::sync::Arc; use anyhow::{Context, Result}; use cudarc::driver::{ - CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtr, DevicePtrMut, LaunchConfig, - PushKernelArg, + CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtr, LaunchConfig, PushKernelArg, }; use ml_core::device::MlDevice; use crate::cfc::snap_features::Mbp10RawInput; use crate::heads::HIDDEN_DIM; -use crate::pinned_mem::{MappedF32Buffer, MappedI32Buffer}; +use crate::pinned_mem::MappedF32Buffer; use rand::{Rng, SeedableRng}; use crate::rl::common::{N_ACTIONS, Q_N_ATOMS, Q_V_MAX, Q_V_MIN}; @@ -1681,117 +1680,86 @@ impl IntegratedTrainer { .forward_logits(h_t_borrow, b_size, &mut pi_logits_d) .context("step_with_lobsim: policy_head.forward_logits")?; - // ── Step 3: read Q logits / V / π logits back to host. ──────── - // Mapped-pinned staging buffers per - // `feedback_no_htod_htoh_only_mapped_pinned`. - let q_logits_host = read_slice_d(&self.stream, &q_logits_d, b_size * k_dqn)?; - let v_pred_host = read_slice_d(&self.stream, &v_pred_d, b_size)?; - let pi_logits_host = read_slice_d(&self.stream, &pi_logits_d, b_size * N_ACTIONS)?; - - // Refresh ISV host mirror so we can read γ for one-step return. - // (`launch_rl_lr_controller` fires inside `step_synthetic` later; - // for the return-bootstrap we only need γ, which the trainer's - // bootstrap path seeds at 0.99 on the first γ-controller emit. - // Until that fire, the mirror reads 0.0 — we fall back to the - // canonical 0.99 default rather than zeroing the bootstrap.) - self.stream - .memcpy_dtoh(&self.isv_d, self.isv_host.as_mut_slice()) - .context("step_with_lobsim: isv dtoh")?; - let gamma = { - let v = self.isv_host[crate::rl::isv_slots::RL_GAMMA_INDEX]; - if v > 0.0 { v } else { 0.99 } - }; - - // ── Step 4: Thompson sample action per batch. ───────────────── - // For each batch, sample one return draw from each action's atom - // distribution and pick the argmax. Bypass `pearl_no_host_branches_ - // in_captured_graph`: this is a host-side sampling step OUTSIDE - // any captured graph; Phase F may move it to a GPU kernel. - let atom_step = (Q_V_MAX - Q_V_MIN) / ((Q_N_ATOMS - 1) as f32); - let atom_supports: Vec = (0..Q_N_ATOMS) - .map(|i| Q_V_MIN + (i as f32) * atom_step) - .collect(); - - // Deterministic per-step RNG seeded from the dqn config seed - // so the toy-bandit test is reproducible per - // `pearl_scoped_init_seed_for_reproducibility`. The salt is the - // host-side step counter (incremented at the end of this method) - // so consecutive calls draw different actions from the same - // Q-distribution without losing reproducibility. - let salt = self.step_counter; + // ── Step 3 (Phase R7b): GPU-pure action sampling. ───────────── + // R4's rl_action_kernel performs Thompson sampling on device + // (per-batch xorshift32 PRNG, per-action softmax + CDF walk, + // argmax over sampled per-action returns). argmax_expected_q + // computes the Bellman-target argmax (deterministic, expected + // value rather than sampled) per + // `pearl_thompson_for_distributional_action_selection`. + // log_pi_at_action computes log π_old via log-softmax + lookup. + // + // The host Thompson + host argmax + host log-softmax loops the + // flawed Phase F+G shipped are GONE. The xorshift32 PRNG state + // is per-batch device-resident; the host-side step_counter + // increment that drove the old ChaCha8 host RNG is also gone. self.step_counter = self.step_counter.wrapping_add(1); - let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64( - self.cfg.dqn_seed.wrapping_add(salt), - ); - let mut actions = vec![0u32; b_size]; - let mut next_actions = vec![0u32; b_size]; - let mut log_pi_old = vec![0.0_f32; b_size]; - - for b in 0..b_size { - // Softmax per-action atom probabilities (apply numerical - // stabilisation per-action row). - let mut action_expected = [0.0_f32; N_ACTIONS]; - let mut sampled_returns = [0.0_f32; N_ACTIONS]; - for a in 0..N_ACTIONS { - let off = b * k_dqn + a * Q_N_ATOMS; - let row = &q_logits_host[off..off + Q_N_ATOMS]; - // Softmax to atom probabilities. - let max_l = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max); - let mut sum = 0.0_f32; - let mut probs = [0.0_f32; Q_N_ATOMS]; - for (i, &l) in row.iter().enumerate() { - probs[i] = (l - max_l).exp(); - sum += probs[i]; - } - if sum > 0.0 { - for p in probs.iter_mut() { - *p /= sum; - } - } - // Expected value of action a — used both for next_actions - // (host-side argmax) and Thompson-sampled-vs-expected - // diagnostics (Phase F). - let mut ev = 0.0_f32; - for i in 0..Q_N_ATOMS { - ev += probs[i] * atom_supports[i]; - } - action_expected[a] = ev; - - // Thompson: sample ONE atom from the categorical, that - // atom's support value is the sampled return. - let u: f32 = rng.gen_range(0.0..1.0); - let mut cum = 0.0_f32; - let mut chosen = Q_N_ATOMS - 1; - for (i, &p) in probs.iter().enumerate() { - cum += p; - if cum >= u { - chosen = i; - break; - } - } - sampled_returns[a] = atom_supports[chosen]; + { + let cfg = LaunchConfig { + grid_dim: (b_size as u32, 1, 1), + block_dim: (N_ACTIONS as u32, 1, 1), + shared_mem_bytes: 0, + }; + let b_size_i = b_size as i32; + let mut launch = self.stream.launch_builder(&self.rl_action_kernel_fn); + launch + .arg(&q_logits_d) + .arg(&self.atom_supports_d) + .arg(&mut self.prng_state_d) + .arg(&mut self.actions_d) + .arg(&b_size_i); + unsafe { + launch + .launch(cfg) + .context("rl_action_kernel launch (Thompson)")?; } + } - // Thompson selector: argmax over sampled per-action returns. - let action_sampled = argmax_f32(&sampled_returns) as u32; - // Bellman target uses argmax of EXPECTED Q (not Thompson sample) - // per the canonical distributional-Q Bellman backup. - let next_action = argmax_f32(&action_expected) as u32; - - // log π_old at the sampled action — softmax of pi_logits row. - let pi_off = b * N_ACTIONS; - let pi_row = &pi_logits_host[pi_off..pi_off + N_ACTIONS]; - let pi_max = pi_row.iter().cloned().fold(f32::NEG_INFINITY, f32::max); - let mut pi_sum = 0.0_f32; - for &l in pi_row.iter() { - pi_sum += (l - pi_max).exp(); + // argmax over expected Q for next_actions (Bellman-target + // action selection — distinct from Thompson per + // `pearl_thompson_for_distributional_action_selection`). + { + let cfg = LaunchConfig { + grid_dim: (b_size as u32, 1, 1), + block_dim: (N_ACTIONS as u32, 1, 1), + shared_mem_bytes: 0, + }; + let b_size_i = b_size as i32; + let mut launch = self.stream.launch_builder(&self.argmax_expected_q_fn); + launch + .arg(&q_logits_d) + .arg(&self.atom_supports_d) + .arg(&mut self.next_actions_d) + .arg(&b_size_i); + unsafe { + launch + .launch(cfg) + .context("argmax_expected_q launch (Bellman target)")?; } - let log_norm = pi_max + pi_sum.ln(); - log_pi_old[b] = pi_row[action_sampled as usize] - log_norm; + } - actions[b] = action_sampled; - next_actions[b] = next_action; + // log π_old at the sampled action — feeds the PPO importance + // ratio. Reads actions_d (just written by rl_action_kernel). + { + let grid_x = ((b_size as u32) + 31) / 32; + let cfg = LaunchConfig { + grid_dim: (grid_x.max(1), 1, 1), + block_dim: (32, 1, 1), + shared_mem_bytes: 0, + }; + let b_size_i = b_size as i32; + let mut launch = self.stream.launch_builder(&self.log_pi_at_action_fn); + launch + .arg(&pi_logits_d) + .arg(&self.actions_d) + .arg(&mut self.log_pi_old_d) + .arg(&b_size_i); + unsafe { + launch + .launch(cfg) + .context("log_pi_at_action launch")?; + } } // ── Step 5 (Phase R6): GPU-pure env step. ───────────────────── @@ -1814,21 +1782,10 @@ impl IntegratedTrainer { .apply_snapshot(last_snap) .context("step_with_lobsim: lobsim.apply_snapshot")?; - // Upload Thompson-sampled actions (computed host-side above - // — Phase R-future moves Thompson to GPU per R4's - // rl_action_kernel; that's a separate refactor since the - // Thompson loop here also computes next_actions + log_pi_old - // by reusing the same softmax pass). - let actions_i32: Vec = actions.iter().map(|&a| a as i32).collect(); - let actions_d = upload_i32( - &self.stream, - actions.as_slice(), - )?; - debug_assert_eq!(actions_i32.len(), b_size); - // GPU translate (action index → market_targets {side, size}). // Reads current position_lots from lobsim.pos_d for Flat-from-* - // actions. + // actions; reads actions from self.actions_d (just written by + // rl_action_kernel above). let pos_bytes_i = lobsim.pos_bytes() as i32; let b_size_i = b_size as i32; { @@ -1841,7 +1798,7 @@ impl IntegratedTrainer { let mut launch = self.stream.launch_builder(&self.actions_to_market_targets_fn); launch - .arg(&actions_d) + .arg(&self.actions_d) .arg(pos_d_ref) .arg(market_targets_d) .arg(&b_size_i) @@ -1887,42 +1844,15 @@ impl IntegratedTrainer { } } - // ── Step 6 (Phase R7a): GPU-pure post-fill pipeline. ────────── - // All advantage/return + reward-scaling + EMA + per-step - // controller launches stay on device. Replaces the host loops - // the flawed Phase F+G shipped per `feedback_cpu_is_read_only`. - - // Upload Thompson-sampled actions + log_pi + next_actions to - // trainer-owned device buffers. R7b lifts Thompson + log_pi + - // next_actions to R4's GPU kernels; the host upload here is - // bounded (3 small buffers of size b_size) and goes away - // entirely after R7b. The u32 → i32 cast matches the kernel - // signature (CUDA `int`). - let actions_i32: Vec = actions.iter().map(|&a| a as i32).collect(); - let next_actions_i32: Vec = next_actions.iter().map(|&a| a as i32).collect(); - self.stream - .memcpy_htod(actions_i32.as_slice(), &mut self.actions_d) - .context("htod actions_d")?; - self.stream - .memcpy_htod(next_actions_i32.as_slice(), &mut self.next_actions_d) - .context("htod next_actions_d")?; - self.stream - .memcpy_htod(log_pi_old.as_slice(), &mut self.log_pi_old_d) - .context("htod log_pi_old_d")?; - - // Upload V(s_t) into a temporary device buffer so - // compute_advantage_return can read it. R-future: keep V(s_t) - // on device throughout (value_head.forward already writes to - // a device buffer; this round-trip exists because v_pred_host - // was already DtoH'd above for the Thompson loop's - // diagnostics). - let v_pred_d = upload_f32(&self.stream, &v_pred_host)?; - // R7a partial: V(s_{t+1}) approximated as V(s_t) until R7b - // wires next_snapshots + forward_encoder(next_snapshots). - // The compute_advantage_return kernel will then produce the - // canonical A_t = r + γV(s_{t+1}) − V(s_t) with the real - // h_{t+1}-derived value. - let v_tp1_d = upload_f32(&self.stream, &v_pred_host)?; + // ── Step 6 (Phase R7b): GPU-pure post-fill pipeline. ────────── + // All actions / next_actions / log_pi_old are already in + // trainer-owned device buffers (R4 kernels wrote them in Step 3). + // No host uploads here. V(s_t) stays in v_pred_d throughout; + // R7c adds next_snapshots + forward_encoder(next_snapshots) + // for true V(s_{t+1}). Until then v_tp1 reuses v_pred_d (the + // h_t approximation — compute_advantage_return still produces + // a valid TD target, just with the bootstrap approximation). + let v_tp1_d_ref = &v_pred_d; // |reward| → reward_abs_d (input for mean_abs_pnl EMA). // Inline launch (vs `launch_abs_copy`) because the per-method @@ -2023,7 +1953,7 @@ impl IntegratedTrainer { .arg(&self.rewards_d) .arg(&self.dones_d) .arg(&v_pred_d) - .arg(&v_tp1_d) + .arg(v_tp1_d_ref) .arg(&mut self.returns_d) .arg(&mut self.advantages_d) .arg(&b_size_i); @@ -2033,7 +1963,10 @@ impl IntegratedTrainer { .context("compute_advantage_return launch")?; } } - let _ = gamma; // gamma is now ISV-resident; host fallback no longer needed here + // γ is read on-device by compute_advantage_return from + // ISV[400] — no host gamma scalar needed in this hot path + // post-R7b (the bootstrap-fallback host read of γ that the + // flawed branch did is gone with the host advantage loop). // ── Step 7: delegate to step_synthetic. ─────────────────────── // step_synthetic now reads trainer-owned device buffers @@ -2279,56 +2212,13 @@ fn accumulate_grad_h( } // ── helpers ────────────────────────────────────────────────────────── - -fn upload_f32(stream: &Arc, host: &[f32]) -> Result> { - let n = host.len(); - let staging = - unsafe { MappedF32Buffer::new(n) }.map_err(|e| anyhow::anyhow!("upload_f32 staging: {e}"))?; - staging.write_from_slice(host); - let mut dst = stream - .alloc_zeros::(n) - .context("upload_f32 alloc")?; - if n > 0 { - let nbytes = n * std::mem::size_of::(); - unsafe { - let (dst_ptr, _g) = dst.device_ptr_mut(stream); - cudarc::driver::result::memcpy_dtod_async( - dst_ptr, - staging.dev_ptr, - nbytes, - stream.cu_stream(), - ) - .context("upload_f32 DtoD")?; - } - } - Ok(dst) -} - -fn upload_i32(stream: &Arc, host: &[u32]) -> Result> { - let n = host.len(); - let mut staging = - unsafe { MappedI32Buffer::new(n) }.map_err(|e| anyhow::anyhow!("upload_i32 staging: {e}"))?; - for (i, &v) in host.iter().enumerate() { - staging.host_slice_mut()[i] = v as i32; - } - let mut dst = stream - .alloc_zeros::(n) - .context("upload_i32 alloc")?; - if n > 0 { - let nbytes = n * std::mem::size_of::(); - unsafe { - let (dst_ptr, _g) = dst.device_ptr_mut(stream); - cudarc::driver::result::memcpy_dtod_async( - dst_ptr, - staging.dev_ptr, - nbytes, - stream.cu_stream(), - ) - .context("upload_i32 DtoD")?; - } - } - Ok(dst) -} +// +// Phase R7b: `upload_f32`, `upload_i32`, and `argmax_f32` are deleted +// because all step_with_lobsim host loops they served are GPU-resident +// now (R3/R4/R5/R6/R7a kernels). Per `feedback_no_hiding` the +// retired helpers are removed, not `#[allow(dead_code)]`'d. The +// remaining `read_*_d` helpers are still load-bearing for the per-step +// loss-scalar readback inside step_synthetic's Q + V backward chain. fn read_scalar_d(stream: &Arc, src: &CudaSlice) -> Result { debug_assert!(src.len() >= 1); @@ -2378,20 +2268,6 @@ fn read_slice_d( Ok(staging.read_all()) } -/// Phase E.3b helper: argmax over a small slice (used for Thompson- -/// selector + Bellman-target next-action). Returns the first index of -/// the maximum value (deterministic tie-break for reproducibility per -/// `pearl_scoped_init_seed_for_reproducibility`). -fn argmax_f32(xs: &[f32]) -> usize { - debug_assert!(!xs.is_empty()); - let mut best_i = 0; - let mut best_v = xs[0]; - for (i, &v) in xs.iter().enumerate().skip(1) { - if v > best_v { - best_v = v; - best_i = i; - } - } - best_i -} +// Phase R7b: `argmax_f32` deleted — R4's `argmax_expected_q` kernel +// is the canonical Bellman-target argmax now.