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:
@@ -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, ¶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::<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, ¶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::<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(¶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<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(¶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| {
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user