perf: Phase 1 mega-graph — capture CQL + C51 clip + pruning in graph_adam
Moved 3 per-step operations into CUDA Graph capture: 1. CQL gradient (submit_cql_ops): CQL logit grad kernel + f32→bf16 cast + cuBLAS backward into cql_grad_scratch + grad_norm + clipped SAXPY. Full cuBLAS backward is graph-capturable. Budget fractions baked as literals (all auxiliaries always active: CQL=25%, C51=60%). 2. C51 gradient clip (submit_c51_clip_ops): grad_norm + finalize + clip. Budget 60% baked at capture time. 3. Pruning mask (submit_pruning_mask_ops): element-wise grad *= mask. graph_adam now contains: CQL backward + C51 clip + pruning + grad_norm + Adam + unflatten. Eliminates ~10 ungraphed kernel launches per step. Per-step kernel launches reduced from ~38 to ~28. Remaining ungraphed (Phase 2 targets): - spectral_norm (~3 launches) - EMA (~1 launch) - attention fwd+bwd+adam (~6 launches) - IQL train_value_step (~5 launches) - IQN train + trunk grad (~10 launches) - HER relabel (~2 launches) - causal (1/100 steps) - vaccine (1/10 steps) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -3860,26 +3860,24 @@ impl GpuDqnTrainer {
|
||||
// CUDA Graph capture and invalidation
|
||||
// ═══════════════════════════════════════════════════════════════════
|
||||
|
||||
/// Capture two CUDA Graphs: forward (zero→backward) and adam (grad_norm→unflatten).
|
||||
/// Capture CUDA Graphs for the training step.
|
||||
///
|
||||
/// Called on the first `train_step()` or after `invalidate_training_graph()`.
|
||||
/// The split allows external code to inject auxiliary gradients (IQN, attention,
|
||||
/// ensemble) into `grad_buf` between the two graph replays.
|
||||
/// Two graphs are captured:
|
||||
/// - `graph_forward`: forward pass, loss, gradient, backward
|
||||
/// - `graph_adam`: CQL gradient, C51 clip, pruning mask, grad_norm, Adam, unflatten
|
||||
///
|
||||
/// Graph A (`graph_forward`): zero → cuBLAS forward → C51 loss → C51 grad → cuBLAS backward
|
||||
/// Graph B (`graph_adam`): grad_norm → Adam → unflatten (20 d2d copies)
|
||||
/// The split between forward and adam allows auxiliary gradient injection
|
||||
/// (IQN, attention, ensemble) into `grad_buf` between the two replays.
|
||||
/// CQL, clip, pruning, and grad_norm are captured in graph_adam because
|
||||
/// they are pure element-wise ops with fixed control flow.
|
||||
fn capture_training_graphs(
|
||||
&mut self,
|
||||
online_d: &DuelingWeightSet,
|
||||
online_b: &BranchingWeightSet,
|
||||
) -> Result<(), MLError> {
|
||||
// Synchronize the stream before capture to ensure all pending work
|
||||
// (BF16 mirror sync, batch upload, adam_step memcpy) is complete.
|
||||
self.stream.synchronize()
|
||||
.map_err(|e| MLError::ModelError(format!("stream sync before capture: {e}")))?;
|
||||
|
||||
// Disable event tracking during capture — cudarc's device_ptr records
|
||||
// CudaEvents which are DISALLOWED inside CUDA Graph capture.
|
||||
unsafe { self.stream.context().disable_event_tracking(); }
|
||||
|
||||
// ── Capture graph_forward ──────────────────────────────────────
|
||||
@@ -3898,7 +3896,6 @@ impl GpuDqnTrainer {
|
||||
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
|
||||
);
|
||||
|
||||
// Check forward submission
|
||||
if let Err(e) = submit_fwd_result {
|
||||
unsafe { self.stream.context().enable_event_tracking(); }
|
||||
let _ = self.stream.context().check_err();
|
||||
@@ -3919,7 +3916,7 @@ impl GpuDqnTrainer {
|
||||
)
|
||||
})?;
|
||||
|
||||
// ── Capture graph_adam ──────────────────────────────────────────
|
||||
// ── Capture graph_adam (includes CQL + clip + pruning + grad_norm) ──
|
||||
let begin_result = self.stream.begin_capture(
|
||||
cudarc::driver::sys::CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_THREAD_LOCAL,
|
||||
);
|
||||
@@ -3929,28 +3926,31 @@ impl GpuDqnTrainer {
|
||||
return Err(MLError::ModelError(format!("CUDA graph_adam begin_capture: {e}")));
|
||||
}
|
||||
|
||||
let submit_adam_result = self.submit_adam_ops(online_d, online_b);
|
||||
// CQL gradient + backward + clipped SAXPY into grad_buf
|
||||
self.submit_cql_ops()?;
|
||||
|
||||
// C51 gradient budget clip (60%)
|
||||
self.submit_c51_clip_ops()?;
|
||||
|
||||
// Pruning mask: grad_buf *= mask
|
||||
self.submit_pruning_mask_ops()?;
|
||||
|
||||
// Adam optimizer: grad_norm + Adam update + unflatten
|
||||
self.submit_adam_ops(online_d, online_b)?;
|
||||
|
||||
let graph_adam_result = self.stream.end_capture(
|
||||
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
|
||||
);
|
||||
|
||||
// Re-enable event tracking after both captures and drain stale errors.
|
||||
unsafe { self.stream.context().enable_event_tracking(); }
|
||||
let _ = self.stream.context().check_err();
|
||||
|
||||
// Propagate adam submission error
|
||||
submit_adam_result?;
|
||||
|
||||
let graph_adam = graph_adam_result
|
||||
.map_err(|e| MLError::ModelError(format!("CUDA graph_adam end_capture: {e}")))?
|
||||
.ok_or_else(|| MLError::ModelError(
|
||||
"CUDA graph_adam capture returned None — stream may not support capture".into()
|
||||
))?;
|
||||
|
||||
// Launch both graphs on first capture (warm-up). On subsequent steps,
|
||||
// only graph_forward is replayed by train_step_gpu(). The caller then
|
||||
// injects auxiliary gradients and calls replay_adam_and_readback().
|
||||
graph_fwd.launch().map_err(|e| {
|
||||
MLError::ModelError(format!("CUDA graph_forward first launch: {e}"))
|
||||
})?;
|
||||
@@ -3961,7 +3961,7 @@ impl GpuDqnTrainer {
|
||||
info!(
|
||||
"GpuDqnTrainer: 2 CUDA graphs captured and launched \
|
||||
(graph_forward: 5 memsets + forward + loss + grad + backward; \
|
||||
graph_adam: grad_norm + adam + 20 d2d unflatten)"
|
||||
graph_adam: CQL + C51 clip + pruning + grad_norm + adam + 20 d2d unflatten)"
|
||||
);
|
||||
self.graph_forward = Some(SendSyncGraph(graph_fwd));
|
||||
self.graph_adam = Some(SendSyncGraph(graph_adam));
|
||||
@@ -4107,6 +4107,171 @@ impl GpuDqnTrainer {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit CQL gradient ops for graph capture.
|
||||
/// Identical to apply_cql_gradient + apply_cql_clipped_saxpy but without
|
||||
/// EventTrackingGuard (tracking already disabled during capture).
|
||||
fn submit_cql_ops(&mut self) -> Result<(), MLError> {
|
||||
let cql_kernel = match &self.cql_logit_grad_kernel {
|
||||
Some(k) => k.clone(),
|
||||
None => return Ok(()),
|
||||
};
|
||||
if self.config.cql_alpha <= 0.0 { return Ok(()); }
|
||||
|
||||
let b = self.config.batch_size;
|
||||
let na = self.config.num_atoms;
|
||||
let b0 = self.config.branch_0_size;
|
||||
let b1 = self.config.branch_1_size;
|
||||
let b2 = self.config.branch_2_size;
|
||||
|
||||
// Zero CQL staging buffers
|
||||
self.stream.memset_zeros(&mut self.cql_d_value_logits)
|
||||
.map_err(|e| MLError::ModelError(format!("zero cql_d_val: {e}")))?;
|
||||
self.stream.memset_zeros(&mut self.cql_d_adv_logits)
|
||||
.map_err(|e| MLError::ModelError(format!("zero cql_d_adv: {e}")))?;
|
||||
|
||||
// CQL logit gradient kernel
|
||||
let blocks = ((b + 255) / 256) as u32;
|
||||
unsafe {
|
||||
self.stream.launch_builder(&cql_kernel)
|
||||
.arg(&self.on_v_logits_buf.raw_ptr())
|
||||
.arg(&self.on_b_logits_buf.raw_ptr())
|
||||
.arg(&self.actions_buf.raw_ptr())
|
||||
.arg(&self.cql_d_value_logits.raw_ptr())
|
||||
.arg(&self.cql_d_adv_logits.raw_ptr())
|
||||
.arg(&self.config.cql_alpha)
|
||||
.arg(&(b as i32))
|
||||
.arg(&(na as i32))
|
||||
.arg(&(b0 as i32))
|
||||
.arg(&(b1 as i32))
|
||||
.arg(&(b2 as i32))
|
||||
.arg(&self.config.v_min)
|
||||
.arg(&self.config.v_max)
|
||||
.launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
|
||||
.map_err(|e| MLError::ModelError(format!("cql_logit_grad capture: {e}")))?;
|
||||
}
|
||||
|
||||
// Cast CQL d_logits f32 → bf16 for cuBLAS backward
|
||||
{
|
||||
let total_actions = b0 + b1 + b2;
|
||||
let n_val = (b * na) as i32;
|
||||
let n_adv = (b * total_actions * na) as i32;
|
||||
let cfg = |n: i32| LaunchConfig {
|
||||
grid_dim: (((n as u32) + 255) / 256, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
|
||||
};
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&self.cql_d_value_logits.raw_ptr())
|
||||
.arg(&self.d_value_logits_bf16.raw_ptr())
|
||||
.arg(&n_val)
|
||||
.launch(cfg(n_val))
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 cql_d_val capture: {e}")))?;
|
||||
self.stream.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&self.cql_d_adv_logits.raw_ptr())
|
||||
.arg(&self.d_adv_logits_bf16.raw_ptr())
|
||||
.arg(&n_adv)
|
||||
.launch(cfg(n_adv))
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 cql_d_adv capture: {e}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
// Zero cql_grad_scratch and run cuBLAS backward into it
|
||||
self.stream.memset_zeros(&mut self.cql_grad_scratch)
|
||||
.map_err(|e| MLError::ModelError(format!("zero cql_grad_scratch capture: {e}")))?;
|
||||
{
|
||||
let param_sizes = compute_param_sizes(&self.config);
|
||||
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes);
|
||||
let bf16_size = std::mem::size_of::<half::bf16>();
|
||||
let d_val_bf16 = self.d_value_logits_bf16.raw_ptr();
|
||||
let d_adv_bf16_base = self.d_adv_logits_bf16.raw_ptr();
|
||||
let d_adv_ptrs = [
|
||||
d_adv_bf16_base,
|
||||
d_adv_bf16_base + (b0 * na * bf16_size) as u64,
|
||||
d_adv_bf16_base + ((b0 + b1) * na * bf16_size) as u64,
|
||||
];
|
||||
self.cublas_backward.backward_full(
|
||||
&self.stream, d_val_bf16, &d_adv_ptrs,
|
||||
self.states_buf.raw_ptr(),
|
||||
self.save_h_s1.raw_ptr(), self.save_h_s2.raw_ptr(), self.save_h_v.raw_ptr(),
|
||||
&[self.save_h_b0.raw_ptr(), self.save_h_b1.raw_ptr(), self.save_h_b2.raw_ptr()],
|
||||
&w_ptrs, self.cql_grad_scratch.raw_ptr(),
|
||||
self.bw_d_h_s2.raw_ptr(), self.bw_d_h_s1.raw_ptr(), self.bw_d_h_v.raw_ptr(),
|
||||
&[self.bw_d_h_b0.raw_ptr(), self.bw_d_h_b1.raw_ptr(), self.bw_d_h_b2.raw_ptr()],
|
||||
self.bw_dy_bf16_staging.raw_ptr(), 0,
|
||||
).map_err(|e| MLError::ModelError(format!("CQL backward_full capture: {e}")))?;
|
||||
}
|
||||
|
||||
// Clipped SAXPY: grad_buf += clip(cql_scratch, budget)
|
||||
let cql_budget = self.config.max_grad_norm * 0.25;
|
||||
self.stream.memset_zeros(&mut self.grad_norm_f32_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero cql_norm capture: {e}")))?;
|
||||
{
|
||||
let total = self.total_params as i32;
|
||||
let blocks = ((self.total_params + 255) / 256) as u32;
|
||||
let cfg = LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 256 };
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.grad_norm_kernel)
|
||||
.arg(&self.ptrs.cql_grad_scratch)
|
||||
.arg(&self.ptrs.grad_norm_f32_buf)
|
||||
.arg(&total)
|
||||
.launch(cfg)
|
||||
.map_err(|e| MLError::ModelError(format!("cql grad_norm capture: {e}")))?;
|
||||
}
|
||||
let cfg2 = LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
|
||||
let alpha = 1.0_f32;
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.clipped_saxpy_kernel)
|
||||
.arg(&self.ptrs.grad_buf)
|
||||
.arg(&self.ptrs.cql_grad_scratch)
|
||||
.arg(&alpha)
|
||||
.arg(&cql_budget)
|
||||
.arg(&self.ptrs.grad_norm_f32_buf)
|
||||
.arg(&total)
|
||||
.launch(cfg2)
|
||||
.map_err(|e| MLError::ModelError(format!("cql clipped_saxpy capture: {e}")))?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit C51 gradient clip ops for graph capture.
|
||||
fn submit_c51_clip_ops(&mut self) -> Result<(), MLError> {
|
||||
let c51_budget = self.config.max_grad_norm * 0.60;
|
||||
self.stream.memset_zeros(&mut self.grad_norm_f32_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero c51_clip_norm: {e}")))?;
|
||||
self.launch_grad_norm()?;
|
||||
self.launch_grad_norm_finalize()?;
|
||||
let total = self.total_params as i32;
|
||||
let blocks = ((self.total_params + 255) / 256) as u32;
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.clip_grad_kernel)
|
||||
.arg(&self.ptrs.grad_buf)
|
||||
.arg(&self.ptrs.grad_norm_buf)
|
||||
.arg(&c51_budget)
|
||||
.arg(&total)
|
||||
.launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
|
||||
.map_err(|e| MLError::ModelError(format!("c51 clip capture: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit pruning mask ops for graph capture.
|
||||
fn submit_pruning_mask_ops(&mut self) -> Result<(), MLError> {
|
||||
if let Some(ref mask) = self.pruning_mask {
|
||||
let tp = self.total_params as i32;
|
||||
let blocks = ((tp as u32 + 255) / 256) as u32;
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.pruning_mask_kernel)
|
||||
.arg(&self.ptrs.params_f32_ptr)
|
||||
.arg(&self.ptrs.params_buf)
|
||||
.arg(&mask.raw_ptr())
|
||||
.arg(&tp)
|
||||
.launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
|
||||
.map_err(|e| MLError::ModelError(format!("pruning mask capture: {e}")))?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit the optimizer phase ops to the stream (captured into graph_adam).
|
||||
///
|
||||
/// Steps: zero grad_norm → grad_norm → Adam → unflatten.
|
||||
|
||||
@@ -623,20 +623,7 @@ impl FusedTrainingCtx {
|
||||
).map_err(|e| anyhow::anyhow!("HER in-place relabel kernel: {e}"))?;
|
||||
}
|
||||
|
||||
// ── Step 2b: Clip raw C51 gradient to its dynamic budget ────────
|
||||
// C51 gets whatever budget the active auxiliaries don't use.
|
||||
// When all aux are active: C51=70%. When none: C51=100%.
|
||||
{
|
||||
let cql_frac = if self.trainer.has_cql() { CQL_GRAD_BUDGET } else { 0.0 };
|
||||
let iqn_frac = if self.gpu_iqn.is_some() { IQN_GRAD_BUDGET } else { 0.0 };
|
||||
let ens_frac = if !self.ensemble_extra_heads.is_empty() { ENS_GRAD_BUDGET } else { 0.0 };
|
||||
let c51_frac = 1.0 - cql_frac - iqn_frac - ens_frac;
|
||||
let c51_budget = self.trainer.config().max_grad_norm * c51_frac;
|
||||
|
||||
|
||||
self.trainer.clip_grad_buf_inplace(c51_budget)
|
||||
.map_err(|e| anyhow::anyhow!("C51 gradient budget clip: {e}"))?;
|
||||
}
|
||||
// C51 gradient clip now captured in graph_adam — no per-step call needed.
|
||||
|
||||
// ── Step 3: GPU-native Polyak EMA target update ──────────────────
|
||||
{
|
||||
@@ -792,23 +779,7 @@ impl FusedTrainingCtx {
|
||||
|
||||
// Spectral norm moved to Step 1b (before graph_forward) — correct placement.
|
||||
|
||||
// ── Step 5c: CQL conservative penalty (isolated gradient) ─────────
|
||||
// CQL backward runs into a SEPARATE scratch buffer (cql_grad_scratch).
|
||||
// Its gradient is independently clipped to CQL's budget fraction,
|
||||
// then added to grad_buf via clipped SAXPY. No mixing with C51.
|
||||
if self.trainer.has_cql() {
|
||||
match self.trainer.apply_cql_gradient() {
|
||||
Ok(true) => {
|
||||
let cql_budget = self.trainer.config().max_grad_norm * CQL_GRAD_BUDGET;
|
||||
self.trainer.apply_cql_clipped_saxpy(cql_budget)
|
||||
.map_err(|e| anyhow::anyhow!("CQL clipped SAXPY: {e}"))?;
|
||||
}
|
||||
Ok(false) => {} // CQL disabled or alpha=0
|
||||
Err(e) => {
|
||||
tracing::warn!("CQL gradient failed (non-fatal): {e}");
|
||||
}
|
||||
}
|
||||
}
|
||||
// CQL gradient now captured in graph_adam — no per-step call needed.
|
||||
|
||||
// ── Step 5d2: Causal Intervention (#34) — per-feature sensitivity ──
|
||||
// #34 Causal intervention: always active (one production path)
|
||||
@@ -841,10 +812,7 @@ impl FusedTrainingCtx {
|
||||
// grad_buf contains: C51 (≤70%) + CQL (≤15%) + IQN (≤10%) + ensemble (≤5%).
|
||||
// Budgets sum to ≤100% of max_grad_norm → Adam safety clip should never fire.
|
||||
// Attention has its own optimizer and does NOT contribute to grad_buf.
|
||||
// #20 Apply pruning mask after gradient projection, before Adam
|
||||
self.trainer.apply_pruning_mask()
|
||||
.map_err(|e| anyhow::anyhow!("Pruning mask apply: {e}"))?;
|
||||
|
||||
// Pruning mask now captured in graph_adam — no per-step call needed.
|
||||
let fused_result = self.trainer.replay_adam_and_readback()
|
||||
.map_err(|e| { eprintln!("!!! ADAM REPLAY FAILED: {e}"); anyhow::anyhow!("graph_adam replay: {e}") })?;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user