perf: migrate Q-stats readback to consolidated pinned buffer
Extend readback_pinned from 3 to 16 floats and write Q-stats DtoH into slots [3..8] instead of the non-pinned q_stats_pending field. On H100 CUDA 13 driver, non-pinned host memory causes cuMemcpyDtoHAsync to degrade to a synchronous copy, creating a latent hang risk. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -629,8 +629,12 @@ pub struct GpuDqnTrainer {
|
||||
|
||||
/// CUDA event for async scalar readback (replaces per-step stream.synchronize)
|
||||
readback_event: Option<CudaEvent>,
|
||||
/// Pinned host buffer [3 × f32] for async DtoH scalar readback.
|
||||
/// Layout: [0]=loss, [1]=mse_loss, [2]=grad_norm_sq.
|
||||
/// Pinned host buffer [16 × f32] for async DtoH scalar readback.
|
||||
/// Layout:
|
||||
/// [0]=loss, [1]=mse_loss, [2]=grad_norm_sq (per-step, from graph_adam)
|
||||
/// [3]=avg_max_q, [4]=q_min, [5]=q_max, [6]=q_mean, [7]=q_var (every 50 steps)
|
||||
/// [8]=causal_mean_sens (every N steps)
|
||||
/// [9..16]=reserved (nan_flags, future use)
|
||||
/// Pinned memory enables true async cuMemcpyDtoHAsync without CPU blocking.
|
||||
readback_pinned: *mut f32,
|
||||
/// Whether there's an in-flight readback to collect
|
||||
@@ -859,8 +863,6 @@ pub struct GpuDqnTrainer {
|
||||
|
||||
/// CUDA event for async Q-stats readback (avoids per-call stream sync)
|
||||
q_stats_event: Option<CudaEvent>,
|
||||
/// Host-side buffer for deferred Q-stats from previous reduction
|
||||
q_stats_pending: [f32; 5],
|
||||
/// Whether there's an in-flight Q-stats readback to collect
|
||||
q_stats_ready: bool,
|
||||
|
||||
@@ -2794,7 +2796,7 @@ impl GpuDqnTrainer {
|
||||
readback_pinned: {
|
||||
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
|
||||
unsafe {
|
||||
cudarc::driver::result::malloc_host(3 * std::mem::size_of::<f32>(), flags)
|
||||
cudarc::driver::result::malloc_host(16 * std::mem::size_of::<f32>(), flags)
|
||||
.map_err(|e| MLError::ModelError(format!("pinned readback alloc: {e}")))?
|
||||
as *mut f32
|
||||
}
|
||||
@@ -2846,7 +2848,6 @@ impl GpuDqnTrainer {
|
||||
q_stats_kernel,
|
||||
q_stats_buf,
|
||||
q_stats_event: None,
|
||||
q_stats_pending: [0.0; 5],
|
||||
q_stats_ready: false,
|
||||
cql_logit_grad_kernel,
|
||||
cql_d_value_logits,
|
||||
@@ -3984,12 +3985,14 @@ impl GpuDqnTrainer {
|
||||
}
|
||||
}
|
||||
self.q_stats_ready = false;
|
||||
QValueStatsResult {
|
||||
avg_max_q: self.q_stats_pending[0] as f64,
|
||||
q_min: self.q_stats_pending[1],
|
||||
q_max: self.q_stats_pending[2],
|
||||
q_mean: self.q_stats_pending[3],
|
||||
q_variance: self.q_stats_pending[4],
|
||||
unsafe {
|
||||
QValueStatsResult {
|
||||
avg_max_q: *self.readback_pinned.add(3) as f64,
|
||||
q_min: *self.readback_pinned.add(4),
|
||||
q_max: *self.readback_pinned.add(5),
|
||||
q_mean: *self.readback_pinned.add(6),
|
||||
q_variance: *self.readback_pinned.add(7),
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// First call — no previous data
|
||||
@@ -4023,10 +4026,10 @@ impl GpuDqnTrainer {
|
||||
.map_err(|e| MLError::ModelError(format!("q_stats_kernel: {e}")))?;
|
||||
}
|
||||
|
||||
// 3. Issue async DtoH for 5 f32s from q_stats_buf
|
||||
// 3. Issue async DtoH for 5 f32s from q_stats_buf into pinned buffer at offset 3
|
||||
unsafe {
|
||||
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
|
||||
self.q_stats_pending.as_mut_ptr().cast(),
|
||||
self.readback_pinned.add(3).cast(),
|
||||
self.q_stats_buf.raw_ptr(),
|
||||
5 * std::mem::size_of::<f32>(),
|
||||
self.stream.cu_stream(),
|
||||
@@ -4064,12 +4067,14 @@ impl GpuDqnTrainer {
|
||||
}
|
||||
}
|
||||
self.q_stats_ready = false;
|
||||
Ok(QValueStatsResult {
|
||||
avg_max_q: self.q_stats_pending[0] as f64,
|
||||
q_min: self.q_stats_pending[1],
|
||||
q_max: self.q_stats_pending[2],
|
||||
q_mean: self.q_stats_pending[3],
|
||||
q_variance: self.q_stats_pending[4],
|
||||
Ok(unsafe {
|
||||
QValueStatsResult {
|
||||
avg_max_q: *self.readback_pinned.add(3) as f64,
|
||||
q_min: *self.readback_pinned.add(4),
|
||||
q_max: *self.readback_pinned.add(5),
|
||||
q_mean: *self.readback_pinned.add(6),
|
||||
q_variance: *self.readback_pinned.add(7),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user