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:
jgrusewski
2026-04-18 11:35:10 +02:00
parent 6960a02af4
commit 69ace546aa

View File

@@ -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()); }
}
}
}