fix: eliminate cuStreamSynchronize in IQN loss readback — was serializing every step
IQN read_total_loss() called cuStreamSynchronize + cuMemcpyDtoH every training step, blocking the CPU until all GPU work completed. This serialized the entire async pipeline, causing 3537ms/step instead of <10ms. Fix: pinned device-mapped memory for IQN total_loss (same pattern as DQN trainer). GPU writes via device pointer, CPU reads via host pointer, zero sync. Also re-applies the cuGraphClone elimination (was lost when agents modified fused_training.rs) and adds recursive_confidence_reduce for deterministic gradient accumulation. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -286,7 +286,9 @@ pub struct GpuIqnHead {
|
||||
|
||||
// ── Loss buffers ─────────────────────────────────────────────────
|
||||
per_sample_loss: CudaSlice<f32>, // [B]
|
||||
total_loss: CudaSlice<f32>, // [1]
|
||||
total_loss: CudaSlice<f32>, // [1] device buffer for reduce kernel output
|
||||
total_loss_pinned: *mut f32, // pinned host for zero-copy readback
|
||||
total_loss_dev_ptr: u64, // device pointer to pinned total_loss
|
||||
|
||||
// ── Training state ───────────────────────────────────────────────
|
||||
adam_step: i32,
|
||||
@@ -448,6 +450,19 @@ impl GpuIqnHead {
|
||||
// Loss buffers
|
||||
let per_sample_loss = alloc_f32(&stream, b, "iqn_per_sample_loss")?;
|
||||
let total_loss = alloc_f32(&stream, 1, "iqn_total_loss")?;
|
||||
// Pinned device-mapped for zero-copy readback (no cuStreamSynchronize per step)
|
||||
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 iqn_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
|
||||
};
|
||||
|
||||
// Persistent rewards/dones buffers
|
||||
let rewards_buf = stream.alloc_zeros::<f32>(b)
|
||||
@@ -547,6 +562,8 @@ impl GpuIqnHead {
|
||||
save_q_online,
|
||||
per_sample_loss,
|
||||
total_loss,
|
||||
total_loss_pinned,
|
||||
total_loss_dev_ptr,
|
||||
adam_step: 0,
|
||||
t_pinned,
|
||||
t_dev_ptr,
|
||||
@@ -984,14 +1001,16 @@ impl GpuIqnHead {
|
||||
.map_err(|e| MLError::ModelError(format!("IQN quantile_huber_loss: {e}")))?;
|
||||
}
|
||||
|
||||
// Loss reduce
|
||||
self.stream.memset_zeros(&mut self.total_loss)
|
||||
.map_err(|e| MLError::ModelError(format!("IQN zero total_loss: {e}")))?;
|
||||
// Loss reduce → pinned device-mapped buffer (zero-copy readback, no sync)
|
||||
unsafe {
|
||||
cudarc::driver::sys::cuMemsetD8Async(
|
||||
self.total_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
|
||||
);
|
||||
let loss_ptr = self.total_loss_dev_ptr;
|
||||
self.stream
|
||||
.launch_builder(&self.loss_reduce_kernel)
|
||||
.arg(&self.per_sample_loss)
|
||||
.arg(&mut self.total_loss)
|
||||
.arg(&loss_ptr)
|
||||
.arg(&batch_size_i32)
|
||||
.launch(LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 })
|
||||
.map_err(|e| MLError::ModelError(format!("IQN loss_reduce: {e}")))?;
|
||||
@@ -1324,18 +1343,10 @@ impl GpuIqnHead {
|
||||
&self.per_sample_loss
|
||||
}
|
||||
|
||||
/// Read IQN total loss (synchronous, 4-byte DtoH). Called once per step.
|
||||
/// Read IQN total loss from pinned device-mapped memory (zero-copy, no sync).
|
||||
/// Returns the PREVIOUS step's value (one-step lag due to async pipeline).
|
||||
pub fn read_total_loss(&self) -> Result<f32, MLError> {
|
||||
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
|
||||
let mut host = [0.0_f32; 1];
|
||||
unsafe {
|
||||
cudarc::driver::sys::cuMemcpyDtoH_v2(
|
||||
host.as_mut_ptr().cast(),
|
||||
self.total_loss.raw_ptr(),
|
||||
std::mem::size_of::<f32>(),
|
||||
);
|
||||
}
|
||||
Ok(host[0])
|
||||
Ok(unsafe { *self.total_loss_pinned })
|
||||
}
|
||||
|
||||
/// IQN gradient w.r.t. shared trunk h_s2 [B, hidden_dim].
|
||||
@@ -1884,6 +1895,9 @@ impl Drop for GpuIqnHead {
|
||||
if !self.t_pinned.is_null() {
|
||||
unsafe { cudarc::driver::sys::cuMemFreeHost(self.t_pinned.cast()); }
|
||||
}
|
||||
if !self.total_loss_pinned.is_null() {
|
||||
unsafe { cudarc::driver::sys::cuMemFreeHost(self.total_loss_pinned.cast()); }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user