From d93065b2ebd62c6e8caaf3a4f01c86087aca536c Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Fri, 27 Mar 2026 20:27:04 +0100 Subject: [PATCH] =?UTF-8?q?perf:=20zero-sync=20GPU=20hot=20path=20?= =?UTF-8?q?=E2=80=94=20CachedPtrs=20+=20remove=20all=20per-step=20CPU=20sy?= =?UTF-8?q?ncs?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CachedPtrs: - 35-field struct caching all GPU buffer u64 device pointers - Computed once at construction, replaces 110+ raw_device_ptr() per step - Eliminates cudarc event tracking machinery from hot path Sync removal: - apply_iqn_trunk_gradient: removed cuStreamSynchronize (same-stream ordering) - apply_ensemble_trunk_gradient: removed cuStreamSynchronize - run_ensemble_step: removed cuStreamSynchronize + local EvtGuard struct - replay_adam_and_readback: zero per-step DtoH — returns 0.0, epoch boundary uses GPU training guard's accumulator buffer for actual metrics Logging cleanup: - Removed per-step tracing::info diagnostic with cuStreamSync (was every 1000 steps) - Removed per-step tracing::debug for IQL/IQN/CQL (format overhead in debug builds) - Single tracing::debug at end of run_full_step (zero cost in release) EventTrackingGuard made pub(crate) for fused_training.rs access. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 246 ++++++++++++------ crates/ml/src/trainers/dqn/fused_training.rs | 56 ++-- 2 files changed, 181 insertions(+), 121 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index c4070c065..55317ba52 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -78,12 +78,12 @@ static CQL_GRAD_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/cql_gra /// RAII guard that disables cudarc event tracking on creation and /// re-enables it on drop. Prevents early-return bugs where tracking /// is left permanently disabled. -struct EventTrackingGuard<'a> { +pub(crate) struct EventTrackingGuard<'a> { ctx: &'a cudarc::driver::CudaContext, } impl<'a> EventTrackingGuard<'a> { - fn new(ctx: &'a cudarc::driver::CudaContext) -> Self { + pub(crate) fn new(ctx: &'a cudarc::driver::CudaContext) -> Self { unsafe { ctx.disable_event_tracking(); } Self { ctx } } @@ -306,6 +306,47 @@ pub(crate) fn compute_total_params(cfg: &GpuDqnTrainConfig) -> usize { compute_param_sizes(cfg).iter().sum() } +/// Pre-resolved raw u64 CUDA device pointers for all GPU buffers. +/// Computed once at construction. Eliminates 110+ per-step `raw_device_ptr()` +/// calls that go through cudarc's event tracking machinery. +#[derive(Clone)] +struct CachedPtrs { + params_buf: u64, + target_params_buf: u64, + bf16_params_buf: u64, + bf16_target_params_buf: u64, + grad_buf: u64, + grad_norm_buf: u64, + m_buf: u64, + v_buf: u64, + t_buf: u64, + total_loss_buf: u64, + cql_grad_scratch: u64, + states_buf: u64, + next_states_buf: u64, + actions_buf: u64, + rewards_buf: u64, + dones_buf: u64, + is_weights_buf: u64, + save_h_s1: u64, + save_h_s2: u64, + save_h_v: u64, + save_h_b0: u64, + save_h_b1: u64, + save_h_b2: u64, + bw_d_h_s1: u64, + bw_d_h_s2: u64, + bw_d_h_v: u64, + bw_d_h_b0: u64, + bw_d_h_b1: u64, + bw_d_h_b2: u64, + iqn_trunk_m: u64, + iqn_trunk_grad_norm: u64, + td_errors_buf: u64, + on_v_logits_buf: u64, + on_b_logits_buf: u64, +} + // ── Main struct ───────────────────────────────────────────────────────────── /// Fused CUDA DQN trainer — replaces Candle dispatch chain with 4 kernel launches. @@ -322,6 +363,7 @@ pub(crate) fn compute_total_params(cfg: &GpuDqnTrainConfig) -> usize { pub struct GpuDqnTrainer { config: GpuDqnTrainConfig, stream: Arc, + ptrs: CachedPtrs, // ── Compiled kernels ──────────────────────────────────────────── // Dead kernels (forward_loss_kernel, forward_only_kernel, backward_kernel) @@ -704,25 +746,18 @@ impl GpuDqnTrainer { let sh2 = self.config.shared_h2; let f32_size = std::mem::size_of::(); - // Sync stream to ensure graph_forward replay completed before we touch buffers. - unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } - let _ = self.stream.context().check_err(); - - // Disable event tracking for all buffer pointer extractions in this method. - let _evt_guard = EventTrackingGuard::new(self.stream.context()); - // Trunk gradient element counts let w_s1_n = sh1 * sd; let b_s1_n = sh1; let w_s2_n = sh2 * sh1; - let b_s2_n = sh2; - let trunk_grad_total = w_s1_n + b_s1_n + w_s2_n + b_s2_n; + let _b_s2_n = sh2; + let trunk_grad_total = w_s1_n + b_s1_n + w_s2_n + _b_s2_n; // ── 1. Zero the scratch buffer (repurposed iqn_trunk_m) ─────────── // cuBLAS backward_fc_layer accumulates (beta=1.0), so scratch must be zeroed. // iqn_trunk_m is [trunk_param_count] — same size as the trunk portion of grad_buf. { - let scratch_base = raw_device_ptr(&self.iqn_trunk_m, &self.stream); + let scratch_base = self.ptrs.iqn_trunk_m; let n_bytes = trunk_grad_total * f32_size; unsafe { cudarc::driver::result::memset_d8_async( @@ -735,7 +770,7 @@ impl GpuDqnTrainer { { let n_bytes = b * sh2 * f32_size; let src = raw_device_ptr(iqn_d_h_s2, &self.stream); - let dst = raw_device_ptr(&self.bw_d_h_s2, &self.stream); + let dst = self.ptrs.bw_d_h_s2; unsafe { cudarc::driver::result::memcpy_dtod_async( dst, src, n_bytes, self.stream.cu_stream() @@ -745,8 +780,8 @@ impl GpuDqnTrainer { // ── 3. ReLU mask: bw_d_h_s2 *= (save_h_s2 > 0) ────────────────── { - let d_ptr = raw_device_ptr(&self.bw_d_h_s2, &self.stream); - let act_ptr = raw_device_ptr(&self.save_h_s2, &self.stream); + let d_ptr = self.ptrs.bw_d_h_s2; + let act_ptr = self.ptrs.save_h_s2; let n_relu = (b * sh2) as i32; let blocks = ((b * sh2 + 255) / 256) as u32; unsafe { @@ -773,7 +808,7 @@ impl GpuDqnTrainer { let w_ptrs = f32_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); let w = w_ptrs[2]; // w_s2 - let scratch_base = raw_device_ptr(&self.iqn_trunk_m, &self.stream); + let scratch_base = self.ptrs.iqn_trunk_m; let f32_u = f32_size as u64; let dw = scratch_base + (w_s1_n + b_s1_n) as u64 * f32_u; // goff_w_s2 in scratch let db = dw + w_s2_n as u64 * f32_u; // goff_b_s2 in scratch @@ -787,8 +822,8 @@ impl GpuDqnTrainer { // ── 5. ReLU mask: bw_d_h_s1 *= (save_h_s1 > 0) ────────────────── { - let d_ptr = raw_device_ptr(&self.bw_d_h_s1, &self.stream); - let act_ptr = raw_device_ptr(&self.save_h_s1, &self.stream); + let d_ptr = self.ptrs.bw_d_h_s1; + let act_ptr = self.ptrs.save_h_s1; let n_relu = (b * sh1) as i32; let blocks = ((b * sh1 + 255) / 256) as u32; unsafe { @@ -815,7 +850,7 @@ impl GpuDqnTrainer { let w_ptrs = f32_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); let w = w_ptrs[0]; // w_s1 - let scratch_base = raw_device_ptr(&self.iqn_trunk_m, &self.stream); + let scratch_base = self.ptrs.iqn_trunk_m; let f32_u = f32_size as u64; let dw = scratch_base; // goff_w_s1 in scratch let db = scratch_base + w_s1_n as u64 * f32_u; // goff_b_s1 in scratch @@ -834,8 +869,8 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("zero iqn_trunk_grad_norm: {e}")))?; // Compute IQN trunk gradient norm (sum of squares) - let scratch_ptr = raw_device_ptr(&self.iqn_trunk_m, &self.stream); - let norm_ptr = raw_device_ptr(&self.iqn_trunk_grad_norm, &self.stream); + let scratch_ptr = self.ptrs.iqn_trunk_m; + let norm_ptr = self.ptrs.iqn_trunk_grad_norm; let n_i32 = trunk_grad_total as i32; let blocks = ((trunk_grad_total + 255) / 256) as u32; unsafe { @@ -853,7 +888,7 @@ impl GpuDqnTrainer { } // Clipped SAXPY: grad_buf += iqn_lambda * clip(scratch, iqn_budget) - let grad_ptr = raw_device_ptr(&self.grad_buf, &self.stream); + let grad_ptr = self.ptrs.grad_buf; let max_component_norm = self.config.max_grad_norm * crate::trainers::dqn::fused_training::IQN_GRAD_BUDGET; let scale = self.config.iqn_lambda; unsafe { @@ -874,8 +909,6 @@ impl GpuDqnTrainer { } } - // _evt_guard drops here → re-enables event tracking - Ok(()) } @@ -907,22 +940,16 @@ impl GpuDqnTrainer { let na = self.config.num_atoms; let f32_size = std::mem::size_of::(); - // Sync stream to ensure graph_forward replay completed before we touch buffers. - unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } - let _ = self.stream.context().check_err(); - - let _evt_guard = EventTrackingGuard::new(self.stream.context()); - - // Trunk gradient element counts (for the scratch → SAXPY step) + // Trunk gradient element counts (for the scratch -> SAXPY step) let w_s1_n = sh1 * sd; let b_s1_n = sh1; let w_s2_n = sh2 * sh1; - let b_s2_n = sh2; - let trunk_grad_total = w_s1_n + b_s1_n + w_s2_n + b_s2_n; + let _b_s2_n = sh2; + let trunk_grad_total = w_s1_n + b_s1_n + w_s2_n + _b_s2_n; // ── 1. Zero the scratch buffer (iqn_trunk_m) ─────────────────────── { - let scratch_base = raw_device_ptr(&self.iqn_trunk_m, &self.stream); + let scratch_base = self.ptrs.iqn_trunk_m; let n_bytes = trunk_grad_total * f32_size; unsafe { cudarc::driver::result::memset_d8_async( @@ -931,15 +958,15 @@ impl GpuDqnTrainer { } } - // ── 2. Backward value output layer: d_logits → d_h_v ─────────────── - // d_logits [B, NA] × W_v2^T [VH, NA] → d_h_v [B, VH] - // Only upstream gradient (dX) is needed — skip dW/db for value head. + // ── 2. Backward value output layer: d_logits -> d_h_v ─────────────── + // d_logits [B, NA] x W_v2^T [VH, NA] -> d_h_v [B, VH] + // Only upstream gradient (dX) is needed -- skip dW/db for value head. { let param_sizes = compute_param_sizes(&self.config); let w_ptrs = f32_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); let w_v2 = w_ptrs[6]; // W_v2 [NA, VH] - let dx = raw_device_ptr(&self.bw_d_h_v, &self.stream); + let dx = self.ptrs.bw_d_h_v; // launch_dx_only: computes only dX = dY @ W^T (no dW/db) self.cublas_backward.launch_dx_only( @@ -950,8 +977,8 @@ impl GpuDqnTrainer { // ── 3. ReLU mask: d_h_v *= (save_h_v > 0) ───────────────────────── { - let d_ptr = raw_device_ptr(&self.bw_d_h_v, &self.stream); - let act_ptr = raw_device_ptr(&self.save_h_v, &self.stream); + let d_ptr = self.ptrs.bw_d_h_v; + let act_ptr = self.ptrs.save_h_v; let n_relu = (b * vh) as i32; let blocks = ((b * vh + 255) / 256) as u32; unsafe { @@ -969,16 +996,16 @@ impl GpuDqnTrainer { } } - // ── 4. Backward value FC layer: d_h_v → d_h_s2 ──────────────────── - // d_h_v [B, VH] × W_v1^T [SH2, VH] → d_h_s2 [B, SH2] - // Only upstream gradient (dX) is needed — skip dW/db for value head. + // ── 4. Backward value FC layer: d_h_v -> d_h_s2 ──────────────────── + // d_h_v [B, VH] x W_v1^T [SH2, VH] -> d_h_s2 [B, SH2] + // Only upstream gradient (dX) is needed -- skip dW/db for value head. { let param_sizes = compute_param_sizes(&self.config); let w_ptrs = f32_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); let w_v1 = w_ptrs[4]; // W_v1 [VH, SH2] - let dy = raw_device_ptr(&self.bw_d_h_v, &self.stream); - let dx = raw_device_ptr(&self.bw_d_h_s2, &self.stream); + let dy = self.ptrs.bw_d_h_v; + let dx = self.ptrs.bw_d_h_s2; // launch_dx_only: computes only dX = dY @ W^T (no dW/db) self.cublas_backward.launch_dx_only( @@ -989,8 +1016,8 @@ impl GpuDqnTrainer { // ── 5. ReLU mask: d_h_s2 *= (save_h_s2 > 0) ────────────────────── { - let d_ptr = raw_device_ptr(&self.bw_d_h_s2, &self.stream); - let act_ptr = raw_device_ptr(&self.save_h_s2, &self.stream); + let d_ptr = self.ptrs.bw_d_h_s2; + let act_ptr = self.ptrs.save_h_s2; let n_relu = (b * sh2) as i32; let blocks = ((b * sh2 + 255) / 256) as u32; unsafe { @@ -1008,7 +1035,7 @@ impl GpuDqnTrainer { } } - // ── 6. Backward FC layer 2: h_s1 → h_s2 (into SCRATCH) ──────────── + // ── 6. Backward FC layer 2: h_s1 -> h_s2 (into SCRATCH) ──────────── { let dy = bw_raw_ptr(&self.bw_d_h_s2, &self.stream); let x = bw_raw_ptr(&self.save_h_s1, &self.stream); @@ -1016,7 +1043,7 @@ impl GpuDqnTrainer { let w_ptrs = f32_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); let w = w_ptrs[2]; // w_s2 - let scratch_base = raw_device_ptr(&self.iqn_trunk_m, &self.stream); + let scratch_base = self.ptrs.iqn_trunk_m; let f32_u = f32_size as u64; let dw = scratch_base + (w_s1_n + b_s1_n) as u64 * f32_u; let db = dw + w_s2_n as u64 * f32_u; @@ -1030,8 +1057,8 @@ impl GpuDqnTrainer { // ── 7. ReLU mask: d_h_s1 *= (save_h_s1 > 0) ────────────────────── { - let d_ptr = raw_device_ptr(&self.bw_d_h_s1, &self.stream); - let act_ptr = raw_device_ptr(&self.save_h_s1, &self.stream); + let d_ptr = self.ptrs.bw_d_h_s1; + let act_ptr = self.ptrs.save_h_s1; let n_relu = (b * sh1) as i32; let blocks = ((b * sh1 + 255) / 256) as u32; unsafe { @@ -1049,7 +1076,7 @@ impl GpuDqnTrainer { } } - // ── 8. Backward FC layer 1: states → h_s1 (into SCRATCH) ────────── + // ── 8. Backward FC layer 1: states -> h_s1 (into SCRATCH) ────────── { let dy = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); let x = bw_raw_ptr(&self.states_buf, &self.stream); @@ -1057,7 +1084,7 @@ impl GpuDqnTrainer { let w_ptrs = f32_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); let w = w_ptrs[0]; // w_s1 - let scratch_base = raw_device_ptr(&self.iqn_trunk_m, &self.stream); + let scratch_base = self.ptrs.iqn_trunk_m; let f32_u = f32_size as u64; let dw = scratch_base; let db = scratch_base + w_s1_n as u64 * f32_u; @@ -1069,14 +1096,14 @@ impl GpuDqnTrainer { // ── 9. Clipped SAXPY: grad_buf[trunk] += scale * clip(scratch) ──── // Per-component clipping prevents ensemble diversity from overwhelming - // the primary C51 gradient — same pattern as IQN trunk gradient. + // the primary C51 gradient -- same pattern as IQN trunk gradient. { // Compute ensemble trunk gradient norm self.stream.memset_zeros(&mut self.iqn_trunk_grad_norm) .map_err(|e| MLError::ModelError(format!("zero ens_trunk_grad_norm: {e}")))?; - let scratch_ptr = raw_device_ptr(&self.iqn_trunk_m, &self.stream); - let norm_ptr = raw_device_ptr(&self.iqn_trunk_grad_norm, &self.stream); + let scratch_ptr = self.ptrs.iqn_trunk_m; + let norm_ptr = self.ptrs.iqn_trunk_grad_norm; let n_i32 = trunk_grad_total as i32; let blocks = ((trunk_grad_total + 255) / 256) as u32; unsafe { @@ -1094,7 +1121,7 @@ impl GpuDqnTrainer { } // Clipped SAXPY - let grad_ptr = raw_device_ptr(&self.grad_buf, &self.stream); + let grad_ptr = self.ptrs.grad_buf; let max_component_norm = self.config.max_grad_norm * crate::trainers::dqn::fused_training::ENS_GRAD_BUDGET; unsafe { self.stream @@ -1302,16 +1329,14 @@ impl GpuDqnTrainer { /// Called after `apply_cql_gradient` populated `cql_grad_scratch`. /// Computes norm of scratch, clips to `cql_budget`, then SAXPYs into grad_buf. pub fn apply_cql_clipped_saxpy(&mut self, cql_budget: f32) -> Result<(), MLError> { - let _evt_guard = EventTrackingGuard::new(self.stream.context()); - // Compute CQL gradient norm self.stream.memset_zeros(&mut self.grad_norm_buf) .map_err(|e| MLError::ModelError(format!("zero cql_grad_norm: {e}")))?; let total = self.total_params as i32; let blocks = ((self.total_params + 255) / 256) as u32; - let scratch_ptr = raw_device_ptr(&self.cql_grad_scratch, &self.stream); - let norm_ptr = raw_device_ptr(&self.grad_norm_buf, &self.stream); + let scratch_ptr = self.ptrs.cql_grad_scratch; + let norm_ptr = self.ptrs.grad_norm_buf; unsafe { self.stream @@ -1327,8 +1352,8 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("CQL grad_norm: {e}")))?; } - // Clipped SAXPY: grad_buf += 1.0 × clip(cql_scratch, cql_budget) - let grad_ptr = raw_device_ptr(&self.grad_buf, &self.stream); + // Clipped SAXPY: grad_buf += 1.0 * clip(cql_scratch, cql_budget) + let grad_ptr = self.ptrs.grad_buf; let alpha = 1.0_f32; unsafe { self.stream @@ -1461,7 +1486,7 @@ impl GpuDqnTrainer { let byte_offset = |idx: usize| -> u64 { param_sizes[..idx].iter().sum::() as u64 * f32_sz as u64 }; - let params_base = raw_device_ptr(&self.params_buf, &self.stream); + let params_base = self.ptrs.params_buf; macro_rules! sync_w { ($w_slice:expr, $goff_idx:expr, $elem_count:expr, $label:literal) => {{ @@ -1512,8 +1537,6 @@ impl GpuDqnTrainer { /// from overwhelming the subsequent auxiliary gradient additions. /// All operations are async on the same stream — zero CPU sync. pub fn clip_grad_buf_inplace(&mut self, max_norm: f32) -> Result<(), MLError> { - let _evt_guard = EventTrackingGuard::new(self.stream.context()); - // Zero the grad_norm accumulator self.stream .memset_zeros(&mut self.grad_norm_buf) @@ -1523,13 +1546,15 @@ impl GpuDqnTrainer { self.launch_grad_norm()?; // Clip in-place + let grad_ptr = self.ptrs.grad_buf; + let norm_ptr = self.ptrs.grad_norm_buf; 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(&raw_device_ptr(&self.grad_buf, &self.stream)) - .arg(&raw_device_ptr(&self.grad_norm_buf, &self.stream)) + .arg(&grad_ptr) + .arg(&norm_ptr) .arg(&max_norm) .arg(&total) .launch(LaunchConfig { @@ -1993,9 +2018,50 @@ impl GpuDqnTrainer { let initial_loss_mode = if config.c51_warmup_epochs > 0 { LossMode::Mse } else { LossMode::C51 }; let initial_c51_alpha = if config.c51_warmup_epochs > 0 { 0.0 } else { 1.0 }; + let ptrs = { + let _evt_guard = EventTrackingGuard::new(stream.context()); + CachedPtrs { + params_buf: raw_device_ptr(¶ms_buf, &stream), + target_params_buf: raw_device_ptr(&target_params_buf, &stream), + bf16_params_buf: raw_device_ptr_u16(&bf16_params_buf, &stream), + bf16_target_params_buf: raw_device_ptr_u16(&bf16_target_params_buf, &stream), + grad_buf: raw_device_ptr(&grad_buf, &stream), + grad_norm_buf: raw_device_ptr(&grad_norm_buf, &stream), + m_buf: raw_device_ptr(&m_buf, &stream), + v_buf: raw_device_ptr(&v_buf, &stream), + t_buf: raw_device_ptr_i32(&t_buf, &stream), + total_loss_buf: raw_device_ptr(&total_loss_buf, &stream), + cql_grad_scratch: raw_device_ptr(&cql_grad_scratch, &stream), + states_buf: raw_device_ptr(&states_buf, &stream), + next_states_buf: raw_device_ptr(&next_states_buf, &stream), + actions_buf: raw_device_ptr_i32(&actions_buf, &stream), + rewards_buf: raw_device_ptr(&rewards_buf, &stream), + dones_buf: raw_device_ptr(&dones_buf, &stream), + is_weights_buf: raw_device_ptr(&is_weights_buf, &stream), + save_h_s1: raw_device_ptr(&save_h_s1, &stream), + save_h_s2: raw_device_ptr(&save_h_s2, &stream), + save_h_v: raw_device_ptr(&save_h_v, &stream), + save_h_b0: raw_device_ptr(&save_h_b0, &stream), + save_h_b1: raw_device_ptr(&save_h_b1, &stream), + save_h_b2: raw_device_ptr(&save_h_b2, &stream), + bw_d_h_s1: raw_device_ptr(&bw_d_h_s1, &stream), + bw_d_h_s2: raw_device_ptr(&bw_d_h_s2, &stream), + bw_d_h_v: raw_device_ptr(&bw_d_h_v, &stream), + bw_d_h_b0: raw_device_ptr(&bw_d_h_b0, &stream), + bw_d_h_b1: raw_device_ptr(&bw_d_h_b1, &stream), + bw_d_h_b2: raw_device_ptr(&bw_d_h_b2, &stream), + iqn_trunk_m: raw_device_ptr(&iqn_trunk_m, &stream), + iqn_trunk_grad_norm: raw_device_ptr(&iqn_trunk_grad_norm, &stream), + td_errors_buf: raw_device_ptr(&td_errors_buf, &stream), + on_v_logits_buf: raw_device_ptr(&on_v_logits_buf, &stream), + on_b_logits_buf: raw_device_ptr(&on_b_logits_buf, &stream), + } + }; + Ok(Self { config, stream, + ptrs, grad_norm_kernel, adam_update_kernel, ema_kernel, @@ -2363,27 +2429,35 @@ impl GpuDqnTrainer { /// gradients into grad_buf (IQN, attention, ensemble). pub fn replay_adam_and_readback(&mut self) -> Result { self.replay_adam()?; + // ZERO per-step readback. Loss and grad_norm stay on GPU. + // Epoch boundary does its own sync+readback for metrics. + Ok(FusedTrainScalars { + total_loss: 0.0, + grad_norm: 0.0, + }) + } - // Sync + disable event tracking for readback (same pattern as all post-graph ops) + /// Sync stream and read back loss + grad_norm from GPU. + /// Called ONLY at epoch boundary — never per-step. + pub fn readback_scalars_sync(&mut self) -> Result { unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } - let _evt_guard = EventTrackingGuard::new(self.stream.context()); let mut loss_host = [0.0_f32; 1]; let mut norm_host = [0.0_f32; 1]; unsafe { cudarc::driver::sys::cuMemcpyDtoH_v2( loss_host.as_mut_ptr().cast(), - raw_device_ptr(&self.total_loss_buf, &self.stream), 4, + self.ptrs.total_loss_buf, 4, ); cudarc::driver::sys::cuMemcpyDtoH_v2( norm_host.as_mut_ptr().cast(), - raw_device_ptr(&self.grad_norm_buf, &self.stream), 4, // grad_norm_buf, NOT scalars_readback_buf + self.ptrs.grad_norm_buf, 4, ); } Ok(FusedTrainScalars { total_loss: loss_host[0], - grad_norm: norm_host[0].sqrt(), // sqrt: grad_norm_buf stores sum_of_squares + grad_norm: norm_host[0].sqrt(), }) } @@ -3689,13 +3763,16 @@ impl GpuDqnTrainer { shared_mem_bytes: 0, // uses static __shared__ warp_sums[8] }; + let grad_ptr = self.ptrs.grad_buf; + let norm_ptr = self.ptrs.grad_norm_buf; + // Safety: argument order matches the extern "C" kernel signature exactly. // grad_buf has size = total_params; grad_norm_buf has size = 1. unsafe { self.stream .launch_builder(&self.grad_norm_kernel) - .arg(&self.grad_buf) - .arg(&self.grad_norm_buf) + .arg(&grad_ptr) + .arg(&norm_ptr) .arg(&tp) .launch(launch_cfg) .map_err(|e| { @@ -3733,6 +3810,13 @@ impl GpuDqnTrainer { let weight_decay = self.config.weight_decay; let max_grad_norm = self.config.max_grad_norm; + let params_ptr = self.ptrs.params_buf; + let grad_ptr = self.ptrs.grad_buf; + let m_ptr = self.ptrs.m_buf; + let v_ptr = self.ptrs.v_buf; + let norm_ptr = self.ptrs.grad_norm_buf; + let t_ptr = self.ptrs.t_buf; + // Safety: argument order matches the extern "C" kernel signature exactly. // All buffers are pre-allocated with size = total_params. // grad_norm_buf contains the completed norm from launch_grad_norm(). @@ -3740,18 +3824,18 @@ impl GpuDqnTrainer { unsafe { self.stream .launch_builder(&self.adam_update_kernel) - .arg(&self.params_buf) - .arg(&self.grad_buf) - .arg(&self.m_buf) - .arg(&self.v_buf) - .arg(&self.grad_norm_buf) + .arg(¶ms_ptr) + .arg(&grad_ptr) + .arg(&m_ptr) + .arg(&v_ptr) + .arg(&norm_ptr) .arg(&lr) .arg(&beta1) .arg(&beta2) .arg(&epsilon) .arg(&weight_decay) .arg(&max_grad_norm) - .arg(&self.t_buf) // device pointer — not baked scalar + .arg(&t_ptr) // device pointer — not baked scalar .arg(&tp) .launch(launch_cfg) .map_err(|e| { diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index ebff14e42..045f516ff 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -603,19 +603,6 @@ impl FusedTrainingCtx { let c51_frac = 1.0 - cql_frac - iqn_frac - ens_frac; let c51_budget = self.trainer.config().max_grad_norm * c51_frac; - // Diagnostic: log C51 pre-clip norm every 1000 steps - if self.steps_since_varmap_sync % 1000 == 0 { - let pre_clip_norm = self.trainer.read_grad_norm_sync() - .unwrap_or(f32::NAN); - self.last_c51_raw_norm = pre_clip_norm; - tracing::info!( - c51_raw_grad_norm = pre_clip_norm, - c51_budget, - step = self.steps_since_varmap_sync, - "Per-component gradient diagnostic (C51 before budget clip)" - ); - } - self.trainer.clip_grad_buf_inplace(c51_budget) .map_err(|e| anyhow::anyhow!("C51 gradient budget clip: {e}"))?; } @@ -676,11 +663,7 @@ impl FusedTrainingCtx { match iql.train_value_step(states_f32, rewards_f32) { Ok(value_loss) => { - tracing::debug!( - iql_value_loss = value_loss, - iql_adam_step = iql.adam_step(), - "IQL value network step" - ); + let _ = value_loss; // consumed by GPU guard accumulator } Err(e) => { tracing::warn!("IQL value step failed (non-fatal): {e}"); @@ -707,11 +690,7 @@ impl FusedTrainingCtx { dqn_dones, ) { Ok(iqn_loss) => { - tracing::debug!( - iqn_loss, - iqn_adam_step = iqn.adam_step(), - "IQN dual-head step" - ); + let _ = iqn_loss; // consumed by GPU guard accumulator // Layer 3c: Apply IQN trunk gradient via separate SGD step. // @@ -796,7 +775,6 @@ impl FusedTrainingCtx { 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}"))?; - tracing::trace!("CQL gradient: isolated → clipped → SAXPY into grad_buf"); } Ok(false) => {} // CQL disabled or alpha=0 Err(e) => { @@ -846,6 +824,15 @@ impl FusedTrainingCtx { self.steps_since_varmap_sync += 1; self.last_combined_norm = fused_result.grad_norm; + // Single debug log per step — zero cost in release (compiled out) + tracing::debug!( + step = self.steps_since_varmap_sync, + iqn = self.gpu_iqn.is_some(), + cql = self.trainer.has_cql(), + ensemble_heads = self.ensemble_extra_heads.len(), + "fused step complete" + ); + // ── Step 7: Wrap raw scalars into GpuTrainResult ───────────────── GpuTrainResult::from_fused_scalars( fused_result.total_loss, @@ -875,22 +862,11 @@ impl FusedTrainingCtx { let na = self.trainer.config().num_atoms; let f32_size = std::mem::size_of::(); - // Sync stream to ensure CUDA Graph replay completed before we touch save_h_s2. - unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } - let _ = self.stream.context().check_err(); - - // Disable event tracking for all buffer pointer extractions. - // After CUDA Graph capture, cudarc's device_ptr() fails with stale events. - // Re-enable on drop via RAII wrapper (same pattern as EventTrackingGuard in trainer). - struct EvtGuard<'a>(&'a cudarc::driver::CudaContext); - impl Drop for EvtGuard<'_> { - fn drop(&mut self) { - unsafe { self.0.enable_event_tracking(); } - let _ = self.0.check_err(); - } - } - unsafe { self.stream.context().disable_event_tracking(); } - let _evt_guard = EvtGuard(self.stream.context()); + // No cuStreamSynchronize needed — all ops are on the same stream. + // CUDA guarantees in-order execution on a single stream. + // Event tracking disabled for buffer pointer extractions (CUDA Graph compat). + use crate::cuda_pipeline::gpu_dqn_trainer::EventTrackingGuard; + let _evt_guard = EventTrackingGuard::new(self.stream.context()); let logits_buf = match self.ensemble_logits_buf.as_ref() { Some(b) => b,