perf: zero-sync GPU hot path — CachedPtrs + remove all per-step CPU syncs

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-27 20:27:04 +01:00
parent e0fe791a50
commit d93065b2eb
2 changed files with 181 additions and 121 deletions

View File

@@ -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<CudaStream>,
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::<f32>();
// 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, &param_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, &param_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::<f32>();
// 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, &param_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, &param_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, &param_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, &param_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::<usize>() 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(&params_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<FusedTrainScalars, MLError> {
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<FusedTrainScalars, MLError> {
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(&params_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| {

View File

@@ -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::<f32>();
// 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,