diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index aa981db26..37127718a 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -705,6 +705,22 @@ pub struct IntegratedTrainer { /// Output of `rl_ensemble_action_value` kernel. Feeds the Q→π /// agreement diagnostic (`rl_q_pi_agree_b`). pub ensemble_q_d: CudaSlice, + + // ── Persistent per-step head output buffers (CUDA Graph stable) ── + // Pre-allocated at init so device pointers are stable across steps, + // enabling CUDA Graph capture of the RL step pipeline. + + /// C51 Q logits at h_t `[B × N_ACTIONS × Q_N_ATOMS]`. + pub q_logits_d: CudaSlice, + /// C51 Q logits at h_{t+1} `[B × N_ACTIONS × Q_N_ATOMS]`. + pub q_logits_tp1_d: CudaSlice, + /// Scalar value prediction at h_t `[B]`. + pub v_pred_d: CudaSlice, + /// Scalar value prediction at h_{t+1} `[B]`. + pub v_pred_tp1_d: CudaSlice, + /// Policy logits at h_t `[B × N_ACTIONS]`. + pub pi_logits_d: CudaSlice, + /// Device-resident xorshift32 PRNG `[B]` for tau sampling. iqn_prng_state_d: CudaSlice, @@ -1488,6 +1504,25 @@ impl IntegratedTrainer { let ensemble_q_d = stream .alloc_zeros::(b_size * N_ACTIONS) .context("alloc ensemble_q_d")?; + + // Persistent per-step head output buffers (CUDA Graph stable). + let k_dqn_alloc = N_ACTIONS * Q_N_ATOMS; + let q_logits_d = stream + .alloc_zeros::(b_size * k_dqn_alloc) + .context("alloc q_logits_d")?; + let q_logits_tp1_d = stream + .alloc_zeros::(b_size * k_dqn_alloc) + .context("alloc q_logits_tp1_d")?; + let v_pred_d = stream + .alloc_zeros::(b_size) + .context("alloc v_pred_d")?; + let v_pred_tp1_d = stream + .alloc_zeros::(b_size) + .context("alloc v_pred_tp1_d")?; + let pi_logits_d = stream + .alloc_zeros::(b_size * N_ACTIONS) + .context("alloc pi_logits_d")?; + // Device-resident PRNG seeds for IQN tau sampling. Starts at // zero; the rl_sample_tau kernel self-seeds on first call using // batch index (no memcpy needed). @@ -1771,6 +1806,11 @@ impl IntegratedTrainer { iqn_q_values_d, iqn_expected_q_d, ensemble_q_d, + q_logits_d, + q_logits_tp1_d, + v_pred_d, + v_pred_tp1_d, + pi_logits_d, iqn_prng_state_d, noisy_exploration, noisy_output_d, @@ -2979,9 +3019,6 @@ impl IntegratedTrainer { // ── Step 3: allocate per-step scratch ──────────────────────── let k_dqn = N_ACTIONS * Q_N_ATOMS; - let mut q_logits_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut pi_logits_d = self.stream.alloc_zeros::(b_size * N_ACTIONS)?; - let mut v_pred_d = self.stream.alloc_zeros::(b_size)?; // Loss accumulators (atomicAdd into single floats from kernels). let q_loss_d = self.stream.alloc_zeros::(1)?; @@ -3041,13 +3078,13 @@ impl IntegratedTrainer { // Step 10 by NOT accumulating q_grad_h_t into the encoder // grad slot). self.dqn_head - .forward(&self.sampled_h_t_d, b_size, &mut q_logits_d) + .forward(&self.sampled_h_t_d, b_size, &mut self.q_logits_d) .context("dqn_head.forward(sampled_h_t) [R7d off-policy]")?; self.policy_head - .forward_logits(h_t_borrow, b_size, &mut pi_logits_d) + .forward_logits(h_t_borrow, b_size, &mut self.pi_logits_d) .context("policy_head.forward_logits")?; self.value_head - .forward(h_t_borrow, &self.isv_d, b_size, &mut v_pred_d) + .forward(h_t_borrow, &self.isv_d, b_size, &mut self.v_pred_d) .context("value_head.forward")?; // R7d: Online Q at sampled_h_tp1 for the Double-DQN argmax that @@ -3055,9 +3092,8 @@ impl IntegratedTrainer { // because online net weights drift faster than transitions // recycle through the replay buffer; storing it at push time // would feed stale-action data into the projection. - let mut q_logits_tp1_sampled_d = self.stream.alloc_zeros::(b_size * k_dqn)?; self.dqn_head - .forward(&self.sampled_h_tp1_d, b_size, &mut q_logits_tp1_sampled_d) + .forward(&self.sampled_h_tp1_d, b_size, &mut self.q_logits_tp1_d) .context("dqn_head.forward(sampled_h_tp1) [R7d Double-DQN argmax src]")?; { let cfg_argmax = LaunchConfig { @@ -3068,7 +3104,7 @@ impl IntegratedTrainer { let b_size_i = b_size as i32; let mut launch = self.stream.launch_builder(&self.argmax_expected_q_fn); launch - .arg(&q_logits_tp1_sampled_d) + .arg(&self.q_logits_tp1_d) .arg(&self.atom_supports_d) .arg(&mut self.sampled_next_actions_d) .arg(&b_size_i); @@ -3086,7 +3122,7 @@ impl IntegratedTrainer { let mut pi_loss_entropy_d_mut = pi_loss_entropy_d; self.policy_head .surrogate_forward( - &pi_logits_d, + &self.pi_logits_d, &self.log_pi_old_d, &self.actions_d, &self.advantages_d, @@ -3141,7 +3177,7 @@ impl IntegratedTrainer { let mut q_loss_d_mut = q_loss_d; self.dqn_head .backward_logits( - &q_logits_d, + &self.q_logits_d, &target_dist_d, &self.sampled_actions_d, b_size, @@ -3181,7 +3217,7 @@ impl IntegratedTrainer { // ── Step 7: PPO backward (logits → grad_w/b/h_t) ───────────── self.policy_head .surrogate_backward_logits( - &pi_logits_d, + &self.pi_logits_d, &self.log_pi_old_d, &self.actions_d, &self.advantages_d, @@ -3196,7 +3232,7 @@ impl IntegratedTrainer { // Couples Q's improved C51 calibration to π's action selection // (per vj5f6 finding that Q was decoupled from policy under // Option B). One block per batch, N_ACTIONS=9 threads per block. - // q_logits_d is still in scope from the earlier forward. + // q_logits_d is a persistent field from the earlier forward. { let cfg_distill = LaunchConfig { grid_dim: (b_size as u32, 1, 1), @@ -3208,8 +3244,8 @@ impl IntegratedTrainer { .stream .launch_builder(&self.rl_q_pi_distill_grad_fn); launch - .arg(&q_logits_d) - .arg(&pi_logits_d) + .arg(&self.q_logits_d) + .arg(&self.pi_logits_d) .arg(&self.atom_supports_d) .arg(&self.isv_d) .arg(&mut pi_grad_logits_d) @@ -3249,7 +3285,7 @@ impl IntegratedTrainer { self.value_head .backward( h_t_borrow, - &v_pred_d, + &self.v_pred_d, &self.returns_d, b_size, &mut v_loss_per_batch_d, @@ -3653,8 +3689,6 @@ impl IntegratedTrainer { // 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)?; @@ -3672,10 +3706,10 @@ impl IntegratedTrainer { // ── 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) + .forward(&self.sampled_h_t_d, b_size, &mut self.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) + .forward(&self.sampled_h_tp1_d, b_size, &mut self.q_logits_tp1_d) .context("dqn_replay_step: dqn_head.forward(sampled_h_tp1)")?; // ── 1b. IQN forward + loss + backward on replay transitions. @@ -3879,7 +3913,7 @@ impl IntegratedTrainer { }; let mut launch = self.stream.launch_builder(&self.argmax_expected_q_fn); launch - .arg(&q_logits_tp1_sampled_d) + .arg(&self.q_logits_tp1_d) .arg(&self.atom_supports_d) .arg(&mut self.sampled_next_actions_d) .arg(&b_size_i); @@ -3918,7 +3952,7 @@ impl IntegratedTrainer { let mut q_loss_d_mut = q_loss_d; self.dqn_head .backward_logits( - &q_logits_d, + &self.q_logits_d, &target_dist_d, &self.sampled_actions_d, b_size, @@ -4123,23 +4157,17 @@ impl IntegratedTrainer { ) .context("frd_head.forward in step_with_lobsim")?; - let k_dqn = N_ACTIONS * Q_N_ATOMS; - let mut q_logits_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_logits_tp1_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut v_pred_d = self.stream.alloc_zeros::(b_size)?; - let mut v_pred_tp1_d = self.stream.alloc_zeros::(b_size)?; - self.dqn_head - .forward(h_t_borrow, b_size, &mut q_logits_d) + .forward(h_t_borrow, b_size, &mut self.q_logits_d) .context("step_with_lobsim: dqn_head.forward(h_t)")?; self.dqn_head - .forward(&self.h_tp1_d, b_size, &mut q_logits_tp1_d) + .forward(&self.h_tp1_d, b_size, &mut self.q_logits_tp1_d) .context("step_with_lobsim: dqn_head.forward(h_tp1) for Double-DQN argmax")?; self.value_head - .forward(h_t_borrow, &self.isv_d, b_size, &mut v_pred_d) + .forward(h_t_borrow, &self.isv_d, b_size, &mut self.v_pred_d) .context("step_with_lobsim: value_head.forward(h_t)")?; self.value_head - .forward(&self.h_tp1_d, &self.isv_d, b_size, &mut v_pred_tp1_d) + .forward(&self.h_tp1_d, &self.isv_d, b_size, &mut self.v_pred_tp1_d) .context("step_with_lobsim: value_head.forward(h_tp1) for true V(s_{t+1})")?; // ── IQN forward alongside C51 ──────────────────────────────── @@ -4196,7 +4224,7 @@ impl IntegratedTrainer { let b_size_i = b_size as i32; let mut launch = self.stream.launch_builder(&self.rl_ensemble_action_value_fn); launch - .arg(&q_logits_d) + .arg(&self.q_logits_d) .arg(&self.atom_supports_d) .arg(&self.iqn_expected_q_d) .arg(&self.isv_d) @@ -4247,9 +4275,8 @@ impl IntegratedTrainer { } // ── Step 2b: Forward π logits for log_pi_old. ───────────────── - let mut pi_logits_d = self.stream.alloc_zeros::(b_size * N_ACTIONS)?; self.policy_head - .forward_logits(h_t_borrow, b_size, &mut pi_logits_d) + .forward_logits(h_t_borrow, b_size, &mut self.pi_logits_d) .context("step_with_lobsim: policy_head.forward_logits")?; // R9 audit — Q-vs-π agreement diag fires into ISV[407] EMA. @@ -4268,8 +4295,8 @@ impl IntegratedTrainer { let b_size_i = b_size as i32; let mut launch = self.stream.launch_builder(&self.rl_q_pi_agree_b_fn); launch - .arg(&q_logits_d) - .arg(&pi_logits_d) + .arg(&self.q_logits_d) + .arg(&self.pi_logits_d) .arg(&self.atom_supports_d) .arg(&self.isv_d) .arg(&b_size_i); @@ -4315,7 +4342,7 @@ impl IntegratedTrainer { let b_size_i = b_size as i32; let mut launch = self.stream.launch_builder(&self.rl_pi_action_kernel_fn); launch - .arg(&pi_logits_d) + .arg(&self.pi_logits_d) .arg(&mut self.prng_state_d) .arg(&mut self.actions_d) .arg(&self.isv_d) @@ -4346,7 +4373,7 @@ impl IntegratedTrainer { let b_size_i = b_size as i32; let mut launch = self.stream.launch_builder(&self.argmax_expected_q_fn); launch - .arg(&q_logits_tp1_d) + .arg(&self.q_logits_tp1_d) .arg(&self.atom_supports_d) .arg(&mut self.next_actions_d) .arg(&b_size_i); @@ -4369,7 +4396,7 @@ impl IntegratedTrainer { 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.pi_logits_d) .arg(&self.actions_d) .arg(&mut self.log_pi_old_d) .arg(&b_size_i); @@ -4397,7 +4424,7 @@ impl IntegratedTrainer { let mut launch = self.stream.launch_builder(&self.rl_confidence_gate_fn); launch .arg(&mut self.actions_d) - .arg(&q_logits_d) + .arg(&self.q_logits_d) .arg(&self.atom_supports_d) .arg(pos_d_ref) .arg(&self.isv_d) @@ -5089,8 +5116,8 @@ impl IntegratedTrainer { .arg(&self.isv_d) .arg(&self.rewards_d) .arg(&self.dones_d) - .arg(&v_pred_d) - .arg(&v_pred_tp1_d) + .arg(&self.v_pred_d) + .arg(&self.v_pred_tp1_d) .arg(&mut self.returns_d) .arg(&mut self.advantages_d) .arg(&b_size_i);