feat(rl): CUDA Graph C capture — replay training step (~25 kernels)

Three-state machine wraps step_synthetic: Q/π/V forwards, Bellman
target projection, C51/PPO losses, backward, reduce_axis0, Adam
updates, FRD backward chain, encoder grad combine + backward, EMA
diagnostic updates, target soft-update. Stream sync + loss readback
(5 mapped-pinned reads) deferred to shared path outside the captured
region.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-25 23:13:31 +02:00
parent 3af2b40100
commit 02987f0111

View File

@@ -427,6 +427,8 @@ pub struct IntegratedTrainer {
postfill_graph: Option<CudaGraph>,
reward_graph: Option<CudaGraph>,
graph_warmup_done: bool,
training_graph: Option<CudaGraph>,
training_warmup_done: bool,
// ── Kernel handles for grad combine + cross-batch reduce ──────────
_grad_h_module: Arc<CudaModule>,
@@ -1905,6 +1907,8 @@ impl IntegratedTrainer {
postfill_graph: None,
reward_graph: None,
graph_warmup_done: false,
training_graph: None,
training_warmup_done: false,
_grad_h_module: grad_h_module,
grad_h_accumulate_fn,
_reduce_axis0_module: reduce_axis0_module,
@@ -3315,6 +3319,28 @@ impl IntegratedTrainer {
let h_t_borrow: &CudaSlice<f32> = self.perception.h_t_view();
debug_assert_eq!(h_t_borrow.len(), b_size * HIDDEN_DIM);
// ── CUDA Graph C: replay training step ───────────────────────
// Three-state machine: first step = warmup (eager), second =
// capture, third+ = replay. All device pointers are stable
// (pre-allocated ss_* fields). Stream sync + loss readback
// stay OUTSIDE the captured region.
if self.training_graph.is_some() {
self.training_graph
.as_ref()
.unwrap()
.launch()
.context("training graph launch")?;
} else {
let capturing_training = self.training_warmup_done;
if !self.training_warmup_done {
self.training_warmup_done = true;
}
if capturing_training {
self.stream
.begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)
.map_err(|e| anyhow::anyhow!("training begin_capture: {e}"))?;
}
// ── Step 3: zero persistent per-step scratch ─────────────────
// Buffers are pre-allocated struct fields for CUDA Graph pointer
// stability. memset_zeros resets them each call.
@@ -3457,9 +3483,7 @@ impl IntegratedTrainer {
&mut self.ss_pi_loss_entropy_d,
)
.context("policy_head.surrogate_forward")?;
let l_pi_host = read_scalar_d(&self.stream, &self.ss_pi_loss_d)?;
let _l_pi_ent_host = read_scalar_d(&self.stream, &self.ss_pi_loss_entropy_d)?;
self.last_pi_loss = l_pi_host;
// Host reads deferred to after graph capture region.
}
// ── Step 6a: Bellman target build (item 2) ───────────────────
@@ -3508,9 +3532,7 @@ impl IntegratedTrainer {
&mut self.ss_q_grad_logits_d,
)
.context("dqn_head.backward_logits (sampled_actions) [R7d off-policy]")?;
let l_q_host = read_scalar_d(&self.stream, &self.ss_q_loss_d)?;
// Mirror for next step's LR controller plateau detection.
self.last_q_loss = l_q_host / (b_size as f32);
// Host read deferred to after graph capture region.
self.dqn_head
.backward_to_w_b_h(
@@ -3615,10 +3637,7 @@ impl IntegratedTrainer {
reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_v_grad_w_per_batch_d, b_size, HIDDEN_DIM, &mut self.ss_v_grad_w_d)?;
reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_v_grad_b_per_batch_d, b_size, 1, &mut self.ss_v_grad_b_d)?;
reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_v_loss_per_batch_d, b_size, 1, &mut self.ss_v_loss_sum_d)?;
let l_v_sum_host = read_scalar_d(&self.stream, &self.ss_v_loss_sum_d)?;
let l_v_host = l_v_sum_host / (b_size as f32);
// Mirror for next step's LR controller plateau detection.
self.last_v_loss = l_v_host;
// Host read deferred to after graph capture region.
// ── Step 9: Adam updates on each head's w and b ──────────────
self.dqn_w_adam
@@ -3649,8 +3668,7 @@ impl IntegratedTrainer {
// Adam steps are no-ops modulo β decay. Once F.5 lands real
// labels, this path activates the supervised learning signal
// without any further trainer-side changes.
let l_frd_host: f32 = {
let frd_n_h = crate::rl::common::FRD_N_HORIZONS;
{
let frd_h_dim = crate::rl::common::FRD_HIDDEN_DIM;
let frd_out_dim = crate::rl::frd::FRD_OUT_DIM;
// Zero FRD scratch buffers.
@@ -3719,11 +3737,8 @@ impl IntegratedTrainer {
self.frd_b2_adam
.step(&mut self.frd_head.b2_d, &self.ss_frd_grad_b2_d)
.context("frd_b2_adam.step")?;
// Per-row CE loss → mean over (b, h). Mapped-pinned read.
let loss_pb_h =
read_slice_d(&self.stream, &self.ss_frd_loss_per_b_h_d, b_size * frd_n_h)?;
loss_pb_h.iter().sum::<f32>() / ((b_size * frd_n_h) as f32)
};
// Host read deferred to after graph capture region.
}
// ── Step 10: encoder grad combine (Phase E.2-DEFER item 1) ───
// Zero the combined slot, fold π+V grad_h_t with λ-weights.
@@ -3917,6 +3932,37 @@ impl IntegratedTrainer {
)
.context("ema_update_per_step(v_grad_norm)")?;
if capturing_training {
let graph = self.stream
.end_capture(
CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
)
.context("training end_capture")?
.ok_or_else(|| anyhow::anyhow!("training end_capture returned None"))?;
self.training_graph = Some(graph);
}
} // end else (training warmup / capture)
// ── Shared path: sync + deferred loss readback ───────────────
// Both the graph-launch path and the warmup/capture path need
// to synchronize before reading mapped-pinned loss scalars.
self.stream.synchronize().context("step_synthetic sync")?;
let l_pi_host = read_scalar_d(&self.stream, &self.ss_pi_loss_d)?;
self.last_pi_loss = l_pi_host;
let l_q_host = read_scalar_d(&self.stream, &self.ss_q_loss_d)?;
self.last_q_loss = l_q_host / (b_size as f32);
let l_v_sum_host = read_scalar_d(&self.stream, &self.ss_v_loss_sum_d)?;
let l_v_host = l_v_sum_host / (b_size as f32);
self.last_v_loss = l_v_host;
let frd_n_h = crate::rl::common::FRD_N_HORIZONS;
let loss_pb_h =
read_slice_d(&self.stream, &self.ss_frd_loss_per_b_h_d, b_size * frd_n_h)?;
let l_frd_host = loss_pb_h.iter().sum::<f32>() / ((b_size * frd_n_h) as f32);
// ── Step 12: compose stats ───────────────────────────────────
// BCE / aux losses are NOT read this phase — perception is driven
// separately by callers via its existing step_batched path. The