diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 0877c8285..cefeb387a 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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::(); + 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. diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index bc227be4c..53c9caba5 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -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}") })?;