fix: eliminate ALL hot-path cuMemcpy — pinned device-mapped everywhere

Per-step (16K/epoch):
- total_loss, mse_loss, grad_norm, q_divergence: CudaSlice → pinned
  device-mapped. GPU kernels write via dev_ptr, CPU reads via host_ptr.
  Zero copies in replay_adam_and_readback (was 4x cuMemcpyDtoHAsync).
- readback_scalars_sync, execute_train_scalars_only: sync DtoH → direct
  pinned read after cuStreamSynchronize.

Per-50-steps:
- eval_v_range: cuMemcpyHtoDAsync → pinned host write (CPU writes
  v_min/v_max, GPU reads via dev_ptr, no copy).
- per_branch_q_gaps: cuMemcpyHtoD → pinned host write (CPU writes 4
  Q-gaps, GPU reads via dev_ptr in qlstm_step + liquid_tau_rk4_step).
- q_stats + q_out readback: stack destination → pinned DtoHAsync
  destination (DMA-capable, faster async transfer).

Structural changes:
- launch_loss_reduce signature: &CudaSlice<f32> → u64 dev_ptr
- loss_gpu_buf/grad_norm_gpu_buf → loss_gpu_ptr/grad_norm_gpu_ptr (u64)
- memset_zeros on CudaSlice → cuMemsetD8Async on dev_ptr
- 6 new pinned allocations in constructor, freed in Drop

Only cuMemcpy remaining: constructor init, checkpoint save/restore,
xavier_init upload, trajectory backtracking, causal intervention,
compute_q_values inference. All per-step training copies eliminated.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-15 12:01:24 +02:00
parent 63dc59c7fb
commit ff562b07b2
3 changed files with 294 additions and 251 deletions

View File

@@ -696,10 +696,15 @@ pub struct GpuDqnTrainer {
// ── Forward output buffers ──────────────────────────────────────
per_sample_loss_buf: CudaSlice<f32>, // [B]
td_errors_buf: CudaSlice<f32>, // [B]
pub(crate) total_loss_buf: CudaSlice<f32>, // [1] float accumulator (deterministic reduce) — C51 loss
pub(crate) mse_loss_buf: CudaSlice<f32>, // [1] MSE loss accumulator (separate from C51)
/// [1] Mean squared Q-divergence between online and target networks (atomicAdd).
pub(crate) q_divergence_buf: CudaSlice<f32>,
/// C51 loss accumulator [1] — pinned device-mapped. GPU writes via dev_ptr, CPU reads via host ptr.
total_loss_pinned: *mut f32,
total_loss_dev_ptr: u64,
/// MSE loss accumulator [1] — pinned device-mapped. GPU writes via dev_ptr, CPU reads via host ptr.
mse_loss_pinned: *mut f32,
mse_loss_dev_ptr: u64,
/// [1] Mean squared Q-divergence — pinned device-mapped. GPU writes via dev_ptr, CPU reads via host ptr.
q_divergence_pinned: *mut f32,
q_divergence_dev_ptr: u64,
// ── Forward-only Q-value output ─────────────────────────────────
q_out_buf: CudaSlice<f32>, // [B, TOTAL_ACTIONS(11)]
@@ -714,7 +719,9 @@ pub struct GpuDqnTrainer {
target_params_buf: CudaSlice<f32>, // [TOTAL_PARAMS] f32 master target parameters (EMA operates here)
m_buf: CudaSlice<f32>, // [TOTAL_PARAMS] Adam first moment (f32 for precision)
v_buf: CudaSlice<f32>, // [TOTAL_PARAMS] Adam second moment (f32 for precision)
pub(crate) grad_norm_buf: CudaSlice<f32>, // [1] f32 L2 norm (written by finalize)
/// Gradient L2 norm [1] — pinned device-mapped. Finalize kernel writes via dev_ptr, CPU reads via host ptr.
grad_norm_pinned: *mut f32,
grad_norm_dev_ptr: u64,
grad_norm_partials: CudaSlice<f32>, // [grad_norm_blocks] per-block partial sums
grad_norm_blocks: usize, // number of blocks for grad_norm kernel
cql_grad_scratch: CudaSlice<f32>, // [TOTAL_PARAMS] f32 CQL gradient isolation buffer
@@ -732,7 +739,8 @@ pub struct GpuDqnTrainer {
/// Per-branch gradient scales [B, 4] — pointer set by IQL trainer.
branch_scales_ptr: u64,
/// Fixed wide v_range [2] for compute_expected_q argmax (not C51 loss).
eval_v_range_buf: CudaSlice<f32>,
/// Pinned device-mapped: CPU writes v_min/v_max, GPU reads via dev_ptr. No HtoD copy.
eval_v_range_pinned: *mut f32,
eval_v_range_ptr: u64,
/// EMA-smoothed Q-stats for stable eval v_range updates.
pub(crate) eval_q_mean_ema: f32,
@@ -1020,6 +1028,9 @@ pub struct GpuDqnTrainer {
q_stats_buf: CudaSlice<f32>,
/// GPU buffer for atom utilization accumulation [2 floats: sum_entropy, sum_utilized]
atom_stats_buf: CudaSlice<f32>,
/// Pinned readback buffer for q_stats [7 floats] + q_out sample 0 [12 floats] = 19 floats.
/// DtoH uses pinned DMA-capable destination for faster async transfer.
q_readback_pinned: *mut f32,
// ── CQL (Conservative Q-Learning) penalty kernel ─────────────────
/// Computes CQL logit gradients: dCQL/d_value_logits and dCQL/d_adv_logits.
@@ -1092,8 +1103,9 @@ pub struct GpuDqnTrainer {
qlstm_n: CudaSlice<f32>,
/// [8] output context vector (written each qlstm_step call)
qlstm_context: CudaSlice<f32>,
/// [4] per-branch Q-gaps device buffer (written in reduce_current_q_stats)
per_branch_q_gaps_buf: CudaSlice<f32>,
/// [4] per-branch Q-gaps — pinned device-mapped. CPU writes, GPU reads via dev_ptr. No HtoD copy.
per_branch_q_gaps_pinned: *mut f32,
per_branch_q_gaps_dev_ptr: u64,
/// qlstm_step kernel handle
qlstm_step_kernel: CudaFunction,
}
@@ -1134,8 +1146,11 @@ impl GpuDqnTrainer {
// Reset Q-divergence EMA — fold 2's divergence baseline differs from fold 1.
self.q_div_ema = 0.0;
self.stream.memset_zeros(&mut self.q_divergence_buf)
.map_err(|e| MLError::ModelError(format!("reset q_divergence: {e}")))?;
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.q_divergence_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
}
tracing::info!("Adam optimizer + PopArt + grad_clip state reset for new fold");
Ok(())
@@ -1214,6 +1229,27 @@ impl Drop for GpuDqnTrainer {
if !self.tau_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.tau_pinned.cast()) };
}
if !self.total_loss_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.total_loss_pinned.cast()) };
}
if !self.mse_loss_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.mse_loss_pinned.cast()) };
}
if !self.grad_norm_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.grad_norm_pinned.cast()) };
}
if !self.eval_v_range_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.eval_v_range_pinned.cast()) };
}
if !self.per_branch_q_gaps_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.per_branch_q_gaps_pinned.cast()) };
}
if !self.q_readback_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.q_readback_pinned.cast()) };
}
if !self.q_divergence_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.q_divergence_pinned.cast()) };
}
}
}
@@ -1225,13 +1261,20 @@ impl GpuDqnTrainer {
pub fn eval_v_range_ptr(&self) -> u64 { self.eval_v_range_ptr }
/// Device pointer to total_loss (pinned device-mapped).
pub fn total_loss_dev_ptr(&self) -> u64 { self.total_loss_dev_ptr }
/// Device pointer to mse_loss (pinned device-mapped).
pub fn mse_loss_dev_ptr(&self) -> u64 { self.mse_loss_dev_ptr }
/// Device pointer to grad_norm (pinned device-mapped).
pub fn grad_norm_dev_ptr(&self) -> u64 { self.grad_norm_dev_ptr }
/// Update per-branch liquid tau modulation — GPU kernel (RK4/Euler adaptive ODE).
///
/// Reads per_branch_q_gaps_buf and qlstm_context from device.
/// Reads per_branch_q_gaps (pinned device-mapped) and qlstm_context from device.
/// Updates per_branch_q_gap_ema_buf and liquid_mod_buf on device.
/// Zero CPU compute.
pub fn update_liquid_tau(&mut self, delta_z_approx: f32) -> Result<(), MLError> {
let gaps_ptr = self.per_branch_q_gaps_buf.raw_ptr();
let gaps_ptr = self.per_branch_q_gaps_dev_ptr;
let ema_ptr = self.per_branch_q_gap_ema_buf.raw_ptr();
let mod_ptr = self.liquid_mod_buf.raw_ptr();
let ctx_ptr = self.qlstm_context.raw_ptr();
@@ -1276,8 +1319,8 @@ impl GpuDqnTrainer {
self.eval_q_std_ema = (1.0 - alpha_std) * self.eval_q_std_ema + alpha_std * q_std;
}
// Keep last_per_branch_q_gaps in sync (per_branch_q_gaps_buf already written in reduce_current_q_stats).
let _ = per_branch_q_gaps; // already on device via per_branch_q_gaps_buf
// Keep last_per_branch_q_gaps in sync (per_branch_q_gaps pinned already written in reduce_current_q_stats).
let _ = per_branch_q_gaps; // already on device via per_branch_q_gaps_dev_ptr
// GPU kernel: liquid tau RK4/Euler adaptive ODE (zero CPU compute).
let delta_z = self.eval_q_std_ema.max(0.01);
@@ -1293,13 +1336,10 @@ impl GpuDqnTrainer {
let half = gap_width;
let v_min = self.eval_q_mean_ema - half;
let v_max = self.eval_q_mean_ema + half;
// Pinned device-mapped: CPU writes directly, GPU reads via dev_ptr. No HtoD copy.
unsafe {
cudarc::driver::sys::cuMemcpyHtoDAsync_v2(
self.eval_v_range_ptr,
[v_min, v_max].as_ptr().cast(),
2 * std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
*self.eval_v_range_pinned = v_min;
*self.eval_v_range_pinned.add(1) = v_max;
}
}
pub fn set_per_sample_support_ptr(&mut self, ptr: u64) {
@@ -1596,9 +1636,9 @@ impl GpuDqnTrainer {
/// Read Q-divergence from pinned readback buffer (offset [12]).
/// Returns the mean squared Q-divergence between online and target networks.
/// Valid after `replay_adam_and_readback()` has been called (async copy landed).
/// Pinned device-mapped: reads directly from host pointer (no DtoH copy).
pub(crate) fn q_divergence_readback(&self) -> f32 {
unsafe { *self.readback_pinned.add(12) }
unsafe { *self.q_divergence_pinned }
}
/// Current IQN readiness scalar [0, 1].
@@ -1612,7 +1652,7 @@ impl GpuDqnTrainer {
/// When divergence is high (networks drifted apart), increase tau to close the gap faster.
/// Uses EMA-smoothed divergence as the baseline.
pub fn compute_adaptive_tau(&mut self, base_tau: f32) -> f32 {
let raw_div = unsafe { *self.readback_pinned.add(12) };
let raw_div = unsafe { *self.q_divergence_pinned };
// EMA smooth the divergence signal
const DIV_BETA: f32 = 0.95;
if self.q_div_ema <= 0.0 {
@@ -2594,16 +2634,8 @@ impl GpuDqnTrainer {
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
// grad_norm_buf[0] now contains the L2 norm directly (sqrt applied by finalize)
let mut norm_val = [0.0_f32; 1];
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
norm_val.as_mut_ptr().cast(),
self.grad_norm_buf.raw_ptr(),
std::mem::size_of::<f32>(),
);
}
Ok(norm_val[0])
// grad_norm is pinned device-mapped — read directly from host pointer after sync.
Ok(unsafe { *self.grad_norm_pinned })
}
/// Compute and return the current grad_buf L2 norm using the STANDALONE
@@ -2622,16 +2654,8 @@ impl GpuDqnTrainer {
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
let mut norm_f32 = [0.0_f32; 1];
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
norm_f32.as_mut_ptr().cast(),
self.grad_norm_buf.raw_ptr(),
std::mem::size_of::<f32>(),
);
}
// grad_norm_buf holds L2 norm directly (finalize computes sqrt)
Ok(norm_f32[0])
// grad_norm is pinned device-mapped — read directly from host pointer after sync.
Ok(unsafe { *self.grad_norm_pinned })
}
/// Apply GPU multi-head feature attention to `save_h_s2` (post-graph).
@@ -2833,12 +2857,44 @@ impl GpuDqnTrainer {
// ── Allocate forward output buffers ─────────────────────────
let per_sample_loss_buf = alloc_f32(&stream, b, "per_sample_loss")?;
let td_errors_buf = alloc_f32(&stream, b, "td_errors")?;
let total_loss_buf = stream.alloc_zeros::<f32>(1)
.map_err(|e| MLError::ModelError(format!("alloc total_loss_f32: {e}")))?;
let mse_loss_buf = stream.alloc_zeros::<f32>(1)
.map_err(|e| MLError::ModelError(format!("alloc mse_loss_f32: {e}")))?;
let q_divergence_buf = stream.alloc_zeros::<f32>(1)
.map_err(|e| MLError::ModelError(format!("alloc q_divergence_f32: {e}")))?;
// total_loss + mse_loss — pinned device-mapped (zero-copy readback).
let total_loss_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned total_loss alloc: {e}")))?
as *mut f32
};
unsafe { *total_loss_pinned = 0.0; }
let total_loss_dev_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp as *mut u64, total_loss_pinned.cast(), 0);
dp
};
let mse_loss_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned mse_loss alloc: {e}")))?
as *mut f32
};
unsafe { *mse_loss_pinned = 0.0; }
let mse_loss_dev_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp as *mut u64, mse_loss_pinned.cast(), 0);
dp
};
// q_divergence — pinned device-mapped (zero-copy readback).
let q_divergence_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned q_divergence alloc: {e}")))?
as *mut f32
};
unsafe { *q_divergence_pinned = 0.0; }
let q_divergence_dev_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp as *mut u64, q_divergence_pinned.cast(), 0);
dp
};
let total_actions = config.branch_0_size + config.branch_1_size + config.branch_2_size + config.branch_3_size;
let q_out_buf = alloc_f32(&stream, b * total_actions, "q_out")?;
let eval_td_snapshot = alloc_f32(&stream, b, "eval_td_snapshot")?;
@@ -2864,7 +2920,19 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("alloc adam_m f32: {e}")))?;
let v_buf = stream.alloc_zeros::<f32>(total_params + cutlass_tile_pad)
.map_err(|e| MLError::ModelError(format!("alloc adam_v f32: {e}")))?;
let grad_norm_buf = alloc_f32(&stream, 1, "grad_norm")?;
// grad_norm — pinned device-mapped (zero-copy readback).
let grad_norm_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned grad_norm alloc: {e}")))?
as *mut f32
};
unsafe { *grad_norm_pinned = 0.0; }
let grad_norm_dev_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp as *mut u64, grad_norm_pinned.cast(), 0);
dp
};
// Size for worst case: d_logits clipping uses val_blocks + adv_blocks
// which can exceed total_params blocks at large batch sizes.
let tba = config.branch_0_size + config.branch_1_size
@@ -2878,14 +2946,22 @@ impl GpuDqnTrainer {
let cql_grad_scratch = stream.alloc_zeros::<f32>(total_params + cutlass_tile_pad)
.map_err(|e| MLError::ModelError(format!("alloc cql_grad_scratch f32: {e}")))?;
// Fixed wide v_range for compute_expected_q (argmax, not C51 loss).
let mut eval_v_range_buf = stream.alloc_zeros::<f32>(2)
.map_err(|e| MLError::ModelError(format!("alloc eval_v_range: {e}")))?;
// Eval v_range: initial from config, updated per-epoch from Q-stats.
let eval_range = [config.v_min, config.v_max];
stream.memcpy_htod(&eval_range, &mut eval_v_range_buf)
.map_err(|e| MLError::ModelError(format!("eval_v_range HtoD: {e}")))?;
let eval_v_range_ptr = eval_v_range_buf.raw_ptr();
// eval_v_range — pinned device-mapped (zero-copy write from CPU).
let eval_v_range_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(2 * std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned eval_v_range alloc: {e}")))?
as *mut f32
};
unsafe {
*eval_v_range_pinned = config.v_min;
*eval_v_range_pinned.add(1) = config.v_max;
}
let eval_v_range_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp as *mut u64, eval_v_range_pinned.cast(), 0);
dp
};
// Adam step counter — pinned device-mapped (no HtoD copies).
let t_pinned: *mut i32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
@@ -3021,6 +3097,14 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("alloc q_stats_f32: {e}")))?;
let atom_stats_buf = stream.alloc_zeros::<f32>(2)
.map_err(|e| MLError::ModelError(format!("alloc atom_stats: {e}")))?;
// q_readback — pinned host buffer for DMA-capable DtoH of q_stats[7] + q_out[12]
let q_readback_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(19 * std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned q_readback alloc: {e}")))?
as *mut f32
};
unsafe { std::ptr::write_bytes(q_readback_pinned, 0, 19); }
info!("GpuDqnTrainer: expected_q + q_stats kernels compiled");
// ── Load mag_concat and strided_accumulate from experience_kernels cubin ──
@@ -3470,13 +3554,13 @@ impl GpuDqnTrainer {
params_ptr: params_buf.raw_ptr(),
target_ptr: target_params_buf.raw_ptr(),
grad_buf: grad_buf.raw_ptr(),
grad_norm_buf: grad_norm_buf.raw_ptr(),
grad_norm_buf: grad_norm_dev_ptr,
m_buf: m_buf.raw_ptr(),
v_buf: v_buf.raw_ptr(),
t_buf: t_dev_ptr,
adaptive_clip_buf: adaptive_clip_dev_ptr,
total_loss_buf: total_loss_buf.raw_ptr(),
mse_loss_buf: mse_loss_buf.raw_ptr(),
total_loss_buf: total_loss_dev_ptr,
mse_loss_buf: mse_loss_dev_ptr,
cql_grad_scratch: cql_grad_scratch.raw_ptr(),
states_buf: states_buf.raw_ptr(),
next_states_buf: next_states_buf.raw_ptr(),
@@ -3694,8 +3778,21 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("alloc qlstm_n: {e}")))?;
let qlstm_context = stream.alloc_zeros::<f32>(QLSTM_HEAD_DIM) // [8]
.map_err(|e| MLError::ModelError(format!("alloc qlstm_context: {e}")))?;
let per_branch_q_gaps_buf = stream.alloc_zeros::<f32>(4) // [4]
.map_err(|e| MLError::ModelError(format!("alloc per_branch_q_gaps_buf: {e}")))?;
// per_branch_q_gaps — pinned device-mapped (zero-copy write from CPU).
let per_branch_q_gaps_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(4 * std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned per_branch_q_gaps alloc: {e}")))?
as *mut f32
};
unsafe {
std::ptr::write_bytes(per_branch_q_gaps_pinned, 0, 4);
}
let per_branch_q_gaps_dev_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp as *mut u64, per_branch_q_gaps_pinned.cast(), 0);
dp
};
let cpbi_module_qlstm = stream.context().load_cubin(EXPECTED_Q_CUBIN.to_vec())
.map_err(|e| MLError::ModelError(format!("cpbi cubin (qlstm): {e}")))?;
let qlstm_step_kernel = cpbi_module_qlstm.load_function("qlstm_step")
@@ -3813,9 +3910,12 @@ impl GpuDqnTrainer {
save_projected,
per_sample_loss_buf,
td_errors_buf,
total_loss_buf,
mse_loss_buf,
q_divergence_buf,
total_loss_pinned,
total_loss_dev_ptr,
mse_loss_pinned,
mse_loss_dev_ptr,
q_divergence_pinned,
q_divergence_dev_ptr,
q_out_buf,
eval_td_snapshot,
eval_loss_snapshot,
@@ -3824,7 +3924,8 @@ impl GpuDqnTrainer {
target_params_buf,
m_buf,
v_buf,
grad_norm_buf,
grad_norm_pinned,
grad_norm_dev_ptr,
grad_norm_partials,
grad_norm_blocks,
cql_grad_scratch,
@@ -3834,7 +3935,7 @@ impl GpuDqnTrainer {
tau_dev_ptr,
per_sample_support_ptr: 0,
branch_scales_ptr: 0,
eval_v_range_buf,
eval_v_range_pinned,
eval_v_range_ptr,
eval_q_mean_ema: 0.0,
eval_q_std_ema: 0.0,
@@ -3918,6 +4019,7 @@ impl GpuDqnTrainer {
q_stats_kernel,
q_stats_buf,
atom_stats_buf,
q_readback_pinned,
cql_logit_grad_kernel,
cql_d_value_logits,
cql_d_adv_logits,
@@ -3996,7 +4098,8 @@ impl GpuDqnTrainer {
qlstm_c,
qlstm_n,
qlstm_context,
per_branch_q_gaps_buf,
per_branch_q_gaps_pinned,
per_branch_q_gaps_dev_ptr,
qlstm_step_kernel,
})
}
@@ -4165,68 +4268,26 @@ impl GpuDqnTrainer {
}
pub fn replay_adam_and_readback(&mut self) -> Result<FusedTrainScalars, MLError> {
// 1. Collect previous step's scalars (if pending)
// Non-blocking: read from pinned buffer without synchronizing.
// These values are monitoring-only (the training guard reads GPU buffers
// directly via loss_gpu_buf()/grad_norm_gpu_buf()). If the async DtoH
// from the previous step hasn't landed yet, we read stale values — acceptable
// for logging scalars. In practice the 12-byte copy completes in <1 µs
// while the GPU does thousands of kernel launches between steps.
let prev_scalars = if self.readback_pending {
self.readback_pending = false;
// Read from pinned buffer: [0]=loss, [1]=mse_loss, [2]=grad_norm
let (loss, mse_loss, grad_norm) = unsafe {
(*self.readback_pinned, *self.readback_pinned.add(1), *self.readback_pinned.add(2))
};
let blended = Self::blend_loss(self.c51_alpha, mse_loss, loss);
FusedTrainScalars {
total_loss: blended,
grad_norm,
}
} else {
// First step — no previous data
FusedTrainScalars { total_loss: 0.0, grad_norm: 0.0 }
// 1. Read previous step's scalars from pinned device-mapped memory.
// These buffers were written by the GPU during the PREVIOUS step's forward
// pass. The current step's forward has not yet overwritten them (it runs
// below). No sync needed — thousands of kernel launches have elapsed since
// the previous forward, guaranteeing the GPU writes have landed in host DRAM.
let (loss, mse_loss, grad_norm) = unsafe {
(*self.total_loss_pinned, *self.mse_loss_pinned, *self.grad_norm_pinned)
};
let blended = Self::blend_loss(self.c51_alpha, mse_loss, loss);
let prev_scalars = FusedTrainScalars {
total_loss: blended,
grad_norm,
};
// 2. Launch current step: grad_norm (two-phase, no atomicAdd) + adam
self.compute_grad_norm_outside_graph()?;
self.replay_adam()?;
// 3. Async DtoH into pinned host buffer (truly async — no CPU blocking)
unsafe {
// [0] = loss
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
self.readback_pinned.cast(),
self.total_loss_buf.raw_ptr(),
std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
// [1] = mse_loss
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
self.readback_pinned.add(1).cast(),
self.mse_loss_buf.raw_ptr(),
std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
// [2] = grad_norm (L2 norm, finalize computes sqrt)
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
self.readback_pinned.add(2).cast(),
self.grad_norm_buf.raw_ptr(),
std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
// [12] = q_divergence (mean squared Q-divergence between online and target)
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
self.readback_pinned.add(12).cast(),
self.q_divergence_buf.raw_ptr(),
std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
}
// No event recording needed — flush_readback reads pinned buffer
// directly without synchronizing (async copy completes in <1 µs).
self.readback_pending = true;
// All per-step scalars are now pinned device-mapped — zero cuMemcpy.
// q_divergence is also pinned: GPU writes via dev_ptr, CPU reads via host ptr.
// Return previous step's values
self.scalars_readback_host = [prev_scalars.total_loss, prev_scalars.grad_norm];
@@ -4242,14 +4303,12 @@ impl GpuDqnTrainer {
/// to have landed. The training guard reads loss/grad_norm from GPU
/// buffers directly — these CPU scalars are monitoring-only.
pub fn flush_readback(&mut self) -> Result<FusedTrainScalars, MLError> {
if !self.readback_pending {
return Ok(FusedTrainScalars { total_loss: 0.0, grad_norm: 0.0 });
}
// Non-blocking: no event.synchronize(). The async DtoH enqueued by
// replay_adam_and_readback() has completed long before epoch end.
self.readback_pending = false;
// Read directly from pinned device-mapped host memory — no DtoH needed.
// The last forward pass wrote these values to host DRAM via PCIe; by epoch
// end, thousands of kernel launches have elapsed, guaranteeing visibility.
let (loss, mse_loss, grad_norm) = unsafe {
(*self.readback_pinned, *self.readback_pinned.add(1), *self.readback_pinned.add(2))
(*self.total_loss_pinned, *self.mse_loss_pinned, *self.grad_norm_pinned)
};
let blended = Self::blend_loss(self.c51_alpha, mse_loss, loss);
Ok(FusedTrainScalars {
@@ -4260,10 +4319,10 @@ impl GpuDqnTrainer {
/// Phase 2: reduce block partials → L2 norm.
/// Single block, 256 threads — reduces grad_norm_partials[0..grad_norm_blocks].
/// Writes L2 norm (with sqrt) into grad_norm_buf.
/// Writes L2 norm (with sqrt) into grad_norm pinned device-mapped buffer.
fn launch_grad_norm_finalize(&self) -> Result<(), MLError> {
let partials_ptr = self.grad_norm_partials.raw_ptr();
let buf_ptr = self.grad_norm_buf.raw_ptr();
let buf_ptr = self.grad_norm_dev_ptr;
let nb = self.grad_norm_blocks as i32;
unsafe {
self.stream
@@ -4289,27 +4348,14 @@ impl GpuDqnTrainer {
unsafe {
cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream());
}
let mut c51_loss_f32 = [0.0_f32; 1];
let mut mse_loss_f32 = [0.0_f32; 1];
let mut norm_val = [0.0_f32; 1];
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
c51_loss_f32.as_mut_ptr().cast(),
self.ptrs.total_loss_buf, std::mem::size_of::<f32>(),
);
cudarc::driver::sys::cuMemcpyDtoH_v2(
mse_loss_f32.as_mut_ptr().cast(),
self.ptrs.mse_loss_buf, std::mem::size_of::<f32>(),
);
cudarc::driver::sys::cuMemcpyDtoH_v2(
norm_val.as_mut_ptr().cast(),
self.ptrs.grad_norm_buf, std::mem::size_of::<f32>(),
);
}
let blended = Self::blend_loss(self.c51_alpha, mse_loss_f32[0], c51_loss_f32[0]);
// After sync, pinned device-mapped host memory is coherent — read directly.
let c51_loss = unsafe { *self.total_loss_pinned };
let mse_loss = unsafe { *self.mse_loss_pinned };
let norm_val = unsafe { *self.grad_norm_pinned };
let blended = Self::blend_loss(self.c51_alpha, mse_loss, c51_loss);
Ok(FusedTrainScalars {
total_loss: blended,
grad_norm: norm_val[0],
grad_norm: norm_val,
})
}
@@ -4344,38 +4390,25 @@ impl GpuDqnTrainer {
self.compute_grad_norm_outside_graph()?;
self.replay_adam()?;
// ── Consolidate readback: gather 4 GPU buffers → 1 readback buf ──
// Layout: [total_loss(1) | mse_loss(1) | grad_norm(1) | td_errors(B)]
// D2D copies are async on the same stream — no host sync until the
// single memcpy_dtoh at the end.
// ── Readback: td_errors from GPU, scalars from pinned host memory ──
// loss/mse_loss/grad_norm are pinned device-mapped — GPU already wrote
// to host DRAM. Only td_errors[B] needs DtoH transfer.
let readback_base = self.readback_buf.raw_ptr();
let f32_bytes = std::mem::size_of::<f32>();
// total_loss (C51) → readback[0]
let loss_src = self.total_loss_buf.raw_ptr();
dtod_copy(readback_base, loss_src, f32_bytes, &self.stream, 0, "readback_gather")?;
// mse_loss → readback[1]
let mse_src = self.mse_loss_buf.raw_ptr();
dtod_copy(readback_base + f32_bytes as u64, mse_src, f32_bytes, &self.stream, 1, "readback_gather")?;
// grad_norm → readback[2]
let norm_src = self.grad_norm_buf.raw_ptr();
dtod_copy(readback_base + (2 * f32_bytes) as u64, norm_src, f32_bytes, &self.stream, 2, "readback_gather")?;
// td_errors → readback[3..3+B]
// td_errors → readback[3..3+B] (only td_errors needs DtoD → DtoH)
let td_src = self.td_errors_buf.raw_ptr();
let td_bytes = b * f32_bytes;
dtod_copy(readback_base + (3 * f32_bytes) as u64, td_src, td_bytes, &self.stream, 3, "readback_gather")?;
// ── Single DtoH transfer ──────────────────────────────────────
// ── Single DtoH transfer (only td_errors portion matters) ─────
super::dtoh_f32(&self.stream, &self.readback_buf, &mut self.readback_host)?;
// ── Unpack on CPU ────────────────────────────────────────────
let c51_loss = self.readback_host[0];
let mse_loss = self.readback_host[1];
let grad_norm_sq = self.readback_host[2];
let td_errors = self.readback_host[3..3 + b].to_vec(); // cpu-side host-buf slice
// ── Unpack: scalars from pinned host, td_errors from readback
let (c51_loss, mse_loss, grad_norm_sq) = unsafe {
(*self.total_loss_pinned, *self.mse_loss_pinned, *self.grad_norm_pinned)
};
let td_errors = self.readback_host[3..3 + b].to_vec();
let blended = Self::blend_loss(self.c51_alpha, mse_loss, c51_loss);
Ok(FusedTrainResult {
@@ -4649,10 +4682,10 @@ impl GpuDqnTrainer {
self.stream.memset_zeros(&mut self.d_adv_logits_mse)
.map_err(|e| MLError::ModelError(format!("vaccine zero mse2: {e}")))?;
self.launch_mse_loss()?;
self.launch_loss_reduce(&self.mse_loss_buf)?;
self.launch_loss_reduce(self.mse_loss_dev_ptr)?;
self.launch_mse_grad_to_scratch()?;
self.launch_c51_loss()?;
self.launch_loss_reduce(&self.total_loss_buf)?;
self.launch_loss_reduce(self.total_loss_dev_ptr)?;
self.launch_c51_grad()?;
// Backward writes g_val to scratch (not grad_buf)
self.launch_cublas_backward_to(g_val_ptr)?;
@@ -4784,29 +4817,15 @@ impl GpuDqnTrainer {
self.compute_grad_norm_outside_graph()?;
self.replay_adam()?;
// ── Raw sync + readback (bypasses stale events from graph capture) ──
// ── Sync then read pinned device-mapped host memory (no DtoH copy) ──
unsafe {
cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream());
}
let mut c51_loss_f32 = [0.0_f32; 1];
let mut mse_loss_f32 = [0.0_f32; 1];
let mut norm_val = [0.0_f32; 1];
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
c51_loss_f32.as_mut_ptr().cast(),
self.total_loss_buf.raw_ptr(), std::mem::size_of::<f32>(),
);
cudarc::driver::sys::cuMemcpyDtoH_v2(
mse_loss_f32.as_mut_ptr().cast(),
self.mse_loss_buf.raw_ptr(), std::mem::size_of::<f32>(),
);
cudarc::driver::sys::cuMemcpyDtoH_v2(
norm_val.as_mut_ptr().cast(),
self.grad_norm_buf.raw_ptr(), std::mem::size_of::<f32>(),
);
}
let blended = Self::blend_loss(self.c51_alpha, mse_loss_f32[0], c51_loss_f32[0]);
self.scalars_readback_host = [blended, norm_val[0]];
let c51_loss = unsafe { *self.total_loss_pinned };
let mse_loss = unsafe { *self.mse_loss_pinned };
let norm_val = unsafe { *self.grad_norm_pinned };
let blended = Self::blend_loss(self.c51_alpha, mse_loss, c51_loss);
self.scalars_readback_host = [blended, norm_val];
Ok(FusedTrainScalars {
total_loss: self.scalars_readback_host[0],
@@ -4857,7 +4876,7 @@ impl GpuDqnTrainer {
// Flag 2: on_b_logits (f32 cuBLAS forward output — branch advantage logits)
self.check_nan_f32(self.on_b_logits_buf.raw_ptr(), b * (b0 + b1 + b2 + b3) * na, 2)?;
// Flag 3: mse_loss_buf [1] (MSE loss scalar)
self.check_nan_f32(self.mse_loss_buf.raw_ptr(), 1, 3)?;
self.check_nan_f32(self.mse_loss_dev_ptr, 1, 3)?;
// Flag 6: grad_buf (cuBLAS backward output) — use ptrs for consistency
self.check_nan_f32(self.ptrs.grad_buf, self.total_params, 6)?;
// Flag 7: save_current_lp (softmax probs from MSE loss — f32)
@@ -5299,26 +5318,36 @@ impl GpuDqnTrainer {
// action selection. Stale values cause 2-state oscillation where
// the system flips between cached results from different steps.
// Cost: ~5µs sync + 28B + 48B DtoH. Runs every 50 training steps.
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
let mut host = [0.0_f32; 7];
// DtoHAsync to pinned readback buffer (DMA-capable destination), then sync.
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
host.as_mut_ptr().cast(),
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
self.q_readback_pinned.cast(),
self.q_stats_buf.raw_ptr(),
7 * std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
}
// Per-branch Q-gap: read first sample's 12 Q-values from q_out_buf.
// Stream already synced. Branch layout: [dir(3), mag(3), ord(3), urg(3)].
let mut q12 = [0.0_f32; 12];
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
q12.as_mut_ptr().cast(),
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
self.q_readback_pinned.add(7).cast(),
self.q_out_buf.raw_ptr(),
12 * std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream());
}
// Read from pinned buffer — host[0..7] = q_stats, host[7..19] = q_out sample 0
let host: [f32; 7] = unsafe {
let mut h = [0.0_f32; 7];
std::ptr::copy_nonoverlapping(self.q_readback_pinned, h.as_mut_ptr(), 7);
h
};
let q12: [f32; 12] = unsafe {
let mut q = [0.0_f32; 12];
std::ptr::copy_nonoverlapping(self.q_readback_pinned.add(7), q.as_mut_ptr(), 12);
q
};
// Per-branch Q-gap: first sample's 12 Q-values.
// Branch layout: [dir(3), mag(3), ord(3), urg(3)].
let branch_sizes = [
self.config.branch_0_size, self.config.branch_1_size,
self.config.branch_2_size, self.config.branch_3_size,
@@ -5333,12 +5362,13 @@ impl GpuDqnTrainer {
offset += bs;
}
// Upload per-branch Q-gaps to device buffer (16 bytes HtoD — once per 50 steps).
// Write per-branch Q-gaps to pinned device-mapped buffer — no HtoD copy.
// GPU reads via dev_ptr (qlstm_step + liquid_tau_rk4_step kernels).
unsafe {
cudarc::driver::sys::cuMemcpyHtoD_v2(
self.per_branch_q_gaps_buf.raw_ptr(),
self.last_per_branch_q_gaps.as_ptr().cast(),
4 * std::mem::size_of::<f32>(),
std::ptr::copy_nonoverlapping(
self.last_per_branch_q_gaps.as_ptr(),
self.per_branch_q_gaps_pinned,
4,
);
}
@@ -5368,13 +5398,13 @@ impl GpuDqnTrainer {
Ok(result)
}
/// Launch the xLSTM mLSTM step: reads q_stats_buf[7] + per_branch_q_gaps_buf[4],
/// Launch the xLSTM mLSTM step: reads q_stats_buf[7] + per_branch_q_gaps[4] (pinned),
/// updates persistent matrix memory C[8,8] and normalizer n[8],
/// writes context[8] output. Runs in a single thread (grid=1, block=1).
/// Called every 50 training steps from reduce_current_q_stats.
pub(crate) fn launch_qlstm_step(&self) -> Result<(), MLError> {
let q_stats_ptr = self.q_stats_buf.raw_ptr();
let gaps_ptr = self.per_branch_q_gaps_buf.raw_ptr();
let gaps_ptr = self.per_branch_q_gaps_dev_ptr;
let c_ptr = self.qlstm_c.raw_ptr();
let n_ptr = self.qlstm_n.raw_ptr();
let ctx_ptr = self.qlstm_context.raw_ptr();
@@ -5557,14 +5587,15 @@ impl GpuDqnTrainer {
/// Pass 3 is submitted separately via `submit_forward_ops_ddqn()`.
pub(crate) fn submit_forward_ops_main(&mut self) -> Result<(), MLError> {
// ── Zero accumulators (all REQUIRED — deterministic reduce / beta=1.0 accumulation) ─
// total_loss_buf: c51_loss_reduce writes this scalar deterministically
self.stream
.memset_zeros(&mut self.total_loss_buf)
.map_err(|e| MLError::ModelError(format!("zero total_loss: {e}")))?;
// mse_loss_buf: c51_loss_reduce writes this scalar deterministically
self.stream
.memset_zeros(&mut self.mse_loss_buf)
.map_err(|e| MLError::ModelError(format!("zero mse_loss: {e}")))?;
// total_loss + mse_loss are pinned device-mapped — zero via cuMemsetD8Async on dev_ptr.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.total_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.mse_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
}
// grad_buf: backward_full uses beta=1.0 GEMM accumulation — zero via ptrs
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
@@ -5594,14 +5625,17 @@ impl GpuDqnTrainer {
self.stream.memset_zeros(&mut self.d_adv_logits_mse)
.map_err(|e| MLError::ModelError(format!("zero d_adv_mse: {e}")))?;
self.launch_mse_loss()?;
self.launch_loss_reduce(&self.mse_loss_buf)?;
self.launch_loss_reduce(self.mse_loss_dev_ptr)?;
self.launch_mse_grad_to_scratch()?;
// C51 path → main buffers (already zeroed above)
self.stream.memset_zeros(&mut self.q_divergence_buf)
.map_err(|e| MLError::ModelError(format!("zero q_divergence: {e}")))?;
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.q_divergence_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
}
self.launch_c51_loss()?;
self.launch_loss_reduce(&self.total_loss_buf)?;
self.launch_loss_reduce(self.total_loss_dev_ptr)?;
self.launch_c51_grad()?;
// Blend: main = α * C51 + (1-α) * MSE
@@ -5713,11 +5747,15 @@ impl GpuDqnTrainer {
/// Submit loss computation + gradient ops (everything between forward and backward).
/// Extracted from submit_forward_ops_main for sub-graph timing.
pub(crate) fn submit_loss_and_grad_ops(&mut self) -> Result<(), MLError> {
// Zero accumulators
self.stream.memset_zeros(&mut self.total_loss_buf)
.map_err(|e| MLError::ModelError(format!("zero total_loss: {e}")))?;
self.stream.memset_zeros(&mut self.mse_loss_buf)
.map_err(|e| MLError::ModelError(format!("zero mse_loss: {e}")))?;
// Zero accumulators — pinned device-mapped, use cuMemsetD8Async on dev_ptr.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.total_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.mse_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
}
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.ptrs.grad_buf,
@@ -5739,14 +5777,17 @@ impl GpuDqnTrainer {
self.stream.memset_zeros(&mut self.d_adv_logits_mse)
.map_err(|e| MLError::ModelError(format!("zero d_adv_mse: {e}")))?;
self.launch_mse_loss()?;
self.launch_loss_reduce(&self.mse_loss_buf)?;
self.launch_loss_reduce(self.mse_loss_dev_ptr)?;
self.launch_mse_grad_to_scratch()?;
// C51 path
self.stream.memset_zeros(&mut self.q_divergence_buf)
.map_err(|e| MLError::ModelError(format!("zero q_divergence: {e}")))?;
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.q_divergence_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
);
}
self.launch_c51_loss()?;
self.launch_loss_reduce(&self.total_loss_buf)?;
self.launch_loss_reduce(self.total_loss_dev_ptr)?;
self.launch_c51_grad()?;
// Blend: main = α * C51 + (1-α) * MSE
@@ -6262,7 +6303,7 @@ impl GpuDqnTrainer {
// ── Outputs (3) ──
.arg(&self.per_sample_loss_buf)
.arg(&self.td_errors_buf)
.arg(&self.total_loss_buf)
.arg(&self.total_loss_dev_ptr)
// ── Saved for backward (2) ──
.arg(&self.save_current_lp)
.arg(&self.save_projected)
@@ -6287,7 +6328,7 @@ impl GpuDqnTrainer {
// ── Spectral decoupling (1) ──
.arg(&self.config.spectral_decoupling_lambda)
// ── Q-divergence accumulator (1) ──
.arg(&self.q_divergence_buf)
.arg(&self.q_divergence_dev_ptr)
// ── Adam step counter for stochastic Expected SARSA ──
.arg(&self.ptrs.t_buf)
.launch(LaunchConfig {
@@ -6303,13 +6344,13 @@ impl GpuDqnTrainer {
/// Deterministic loss reduction: sequential sum of per_sample_loss → total_loss / batch_size.
/// Grid=(1,1,1), Block=(1,1,1). Replaces atomicAdd for fully deterministic training.
fn launch_loss_reduce(&self, total_loss_buf: &CudaSlice<f32>) -> Result<(), MLError> {
fn launch_loss_reduce(&self, total_loss_dev_ptr: u64) -> Result<(), MLError> {
let b = self.config.batch_size as i32;
unsafe {
self.stream
.launch_builder(&self.c51_loss_reduce_kernel)
.arg(&self.per_sample_loss_buf)
.arg(total_loss_buf)
.arg(&total_loss_dev_ptr)
.arg(&b)
.launch(LaunchConfig {
grid_dim: (1, 1, 1),
@@ -6464,7 +6505,7 @@ impl GpuDqnTrainer {
// ── Outputs (3) ──
.arg(&self.per_sample_loss_buf)
.arg(&self.td_errors_buf)
.arg(&self.mse_loss_buf) // MSE writes to separate accumulator (not total_loss_buf)
.arg(&self.mse_loss_dev_ptr) // MSE writes to separate accumulator (not total_loss)
// ── Saved for backward (2) ── repurposed: save_current_lp = softmax probs,
// save_projected = per-branch E[Q] values (4 floats per sample per branch)
.arg(&self.save_current_lp)

View File

@@ -2268,18 +2268,20 @@ impl FusedTrainingCtx {
Ok(())
}
pub(crate) fn loss_gpu_buf(&self) -> &cudarc::driver::CudaSlice<f32> {
/// Raw device pointer to the active loss scalar (pinned device-mapped).
/// Returns MSE dev_ptr when c51_alpha ≈ 0, otherwise C51 total_loss dev_ptr.
pub(crate) fn loss_gpu_ptr(&self) -> u64 {
if self.trainer.c51_alpha() < 1e-6 {
&self.trainer.mse_loss_buf
self.trainer.mse_loss_dev_ptr()
} else {
&self.trainer.total_loss_buf
self.trainer.total_loss_dev_ptr()
}
}
/// Return a reference to the GPU-resident grad_norm scalar (f32).
/// Raw device pointer to grad_norm scalar (pinned device-mapped).
/// Written by the CUDA graph's grad_norm kernel — valid after `replay_adam()`.
pub(crate) fn grad_norm_gpu_buf(&self) -> &cudarc::driver::CudaSlice<f32> {
&self.trainer.grad_norm_buf
pub(crate) fn grad_norm_gpu_ptr(&self) -> u64 {
self.trainer.grad_norm_dev_ptr()
}
/// Read NaN detection flags (synchronizes stream). Used by training guard on NaN halt.

View File

@@ -1376,8 +1376,8 @@ impl DQNTrainer {
if train_step_count % 5 == 0 {
if let Some(ref mut guard) = self.training_guard {
let fused = self.fused_ctx.as_ref().unwrap();
let loss_raw = fused.loss_gpu_buf().raw_ptr();
let grad_raw = fused.grad_norm_gpu_buf().raw_ptr();
let loss_raw = fused.loss_gpu_ptr();
let grad_raw = fused.grad_norm_gpu_ptr();
let gr = guard.check_and_accumulate(
loss_raw,
grad_raw,