refactor(rl): pre-allocate per-step head output buffers for graph capture

Move q_logits_d, q_logits_tp1_d, v_pred_d, v_pred_tp1_d, pi_logits_d
from per-step alloc_zeros to persistent trainer fields allocated at
init. Stable device pointers are a prerequisite for CUDA Graph capture.

Updated step_with_lobsim, step_synthetic, and dqn_replay_step to use
self.field instead of local allocations.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-25 22:11:56 +02:00
parent a548adf7b7
commit 0b7895a4fb

View File

@@ -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<f32>,
// ── 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<f32>,
/// C51 Q logits at h_{t+1} `[B × N_ACTIONS × Q_N_ATOMS]`.
pub q_logits_tp1_d: CudaSlice<f32>,
/// Scalar value prediction at h_t `[B]`.
pub v_pred_d: CudaSlice<f32>,
/// Scalar value prediction at h_{t+1} `[B]`.
pub v_pred_tp1_d: CudaSlice<f32>,
/// Policy logits at h_t `[B × N_ACTIONS]`.
pub pi_logits_d: CudaSlice<f32>,
/// Device-resident xorshift32 PRNG `[B]` for tau sampling.
iqn_prng_state_d: CudaSlice<u32>,
@@ -1488,6 +1504,25 @@ impl IntegratedTrainer {
let ensemble_q_d = stream
.alloc_zeros::<f32>(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::<f32>(b_size * k_dqn_alloc)
.context("alloc q_logits_d")?;
let q_logits_tp1_d = stream
.alloc_zeros::<f32>(b_size * k_dqn_alloc)
.context("alloc q_logits_tp1_d")?;
let v_pred_d = stream
.alloc_zeros::<f32>(b_size)
.context("alloc v_pred_d")?;
let v_pred_tp1_d = stream
.alloc_zeros::<f32>(b_size)
.context("alloc v_pred_tp1_d")?;
let pi_logits_d = stream
.alloc_zeros::<f32>(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::<f32>(b_size * k_dqn)?;
let mut pi_logits_d = self.stream.alloc_zeros::<f32>(b_size * N_ACTIONS)?;
let mut v_pred_d = self.stream.alloc_zeros::<f32>(b_size)?;
// Loss accumulators (atomicAdd into single floats from kernels).
let q_loss_d = self.stream.alloc_zeros::<f32>(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::<f32>(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::<f32>(b_size * k_dqn)?;
let mut q_logits_tp1_sampled_d = self.stream.alloc_zeros::<f32>(b_size * k_dqn)?;
let q_loss_d = self.stream.alloc_zeros::<f32>(1)?;
let mut q_target_full_d = self.stream.alloc_zeros::<f32>(b_size * k_dqn)?;
let mut q_target_action_d = self.stream.alloc_zeros::<f32>(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::<f32>(b_size * k_dqn)?;
let mut q_logits_tp1_d = self.stream.alloc_zeros::<f32>(b_size * k_dqn)?;
let mut v_pred_d = self.stream.alloc_zeros::<f32>(b_size)?;
let mut v_pred_tp1_d = self.stream.alloc_zeros::<f32>(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::<f32>(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);