perf: cuMemsetD8Async → cuMemsetD32Async for 4x memset bandwidth

All gradient buffer clears now use cuMemsetD32Async (u32-wide writes)
instead of cuMemsetD8Async (byte-wide). 4x memory bandwidth utilization
for the ~33 memset calls per training step. Size params converted from
bytes to f32 element count (.num_bytes() → .len()).

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-19 14:15:45 +02:00
parent 951dd97d71
commit a79b8761cf
4 changed files with 57 additions and 57 deletions

View File

@@ -639,8 +639,8 @@ impl GpuAttention {
// Zero grad_buf — raw cuMemsetD8Async, no cudarc overhead during graph capture
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.d_params.raw_ptr(), 0, self.d_params.num_bytes(),
cudarc::driver::sys::cuMemsetD32Async(
self.d_params.raw_ptr(), 0, self.d_params.len(),
stream.cu_stream(),
);
}

View File

@@ -1649,8 +1649,8 @@ impl GpuDqnTrainer {
// Reset Q-divergence EMA — fold 2's divergence baseline differs from fold 1.
self.q_div_ema = 0.0;
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.q_divergence_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
self.q_divergence_dev_ptr, 0, 1, self.stream.cu_stream(),
);
}
@@ -1673,8 +1673,8 @@ impl GpuDqnTrainer {
let m_ptr = self.m_buf.raw_ptr() + start_byte;
let v_ptr = self.v_buf.raw_ptr() + start_byte;
unsafe {
cudarc::driver::sys::cuMemsetD8Async(m_ptr, 0, range_bytes, self.stream.cu_stream());
cudarc::driver::sys::cuMemsetD8Async(v_ptr, 0, range_bytes, self.stream.cu_stream());
cudarc::driver::sys::cuMemsetD32Async(m_ptr, 0, range_bytes / 4, self.stream.cu_stream());
cudarc::driver::sys::cuMemsetD32Async(v_ptr, 0, range_bytes / 4, self.stream.cu_stream());
}
Ok(())
}
@@ -1832,8 +1832,8 @@ impl GpuDqnTrainer {
let homeostatic_total_buf_ptr = self.homeostatic_total_buf.raw_ptr();
let homeostatic_penalties_buf_ptr = self.homeostatic_penalties_buf.raw_ptr();
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
homeostatic_total_buf_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
homeostatic_total_buf_ptr, 0, 1, self.stream.cu_stream(),
);
}
unsafe {
@@ -2863,10 +2863,10 @@ impl GpuDqnTrainer {
let save_h_b3_ptr = self.save_h_b3.raw_ptr();
// Use raw memset to avoid &mut self borrow conflict on the penalty buf
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
branch_indep_penalty_buf_ptr,
0,
std::mem::size_of::<f32>(),
1,
self.stream.cu_stream(),
);
}
@@ -2899,10 +2899,10 @@ impl GpuDqnTrainer {
let q_out_buf_ptr = self.q_out_buf.raw_ptr();
// Use raw memset to avoid &mut self borrow conflict on the penalty buf
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
temporal_penalty_buf_ptr,
0,
std::mem::size_of::<f32>(),
1,
self.stream.cu_stream(),
);
}
@@ -2929,7 +2929,7 @@ impl GpuDqnTrainer {
/// signal for the trunk. Uses save_h_s2 (enriched after mamba2_step).
///
/// Graph-safe: writes per-sample loss to predictive_per_sample_buf, then
/// reduces via c51_loss_reduce kernel. No cuMemsetD8Async or atomicAdd.
/// reduces via c51_loss_reduce kernel. No cuMemsetD32Async or atomicAdd.
pub(crate) fn compute_predictive_coding_loss(&self, batch_size: usize) -> Result<(), MLError> {
let sh2 = self.config.shared_h2 as i32;
let save_h_s2_ptr = self.save_h_s2.raw_ptr();
@@ -4118,11 +4118,11 @@ impl GpuDqnTrainer {
let scratch_d_h_b3 = self.bw_d_h_b3.raw_ptr();
// Zero CQL scratch buffer (backward_full uses beta=1.0 accumulation)
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
// Raw cuMemsetD32Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.cql_grad_scratch.raw_ptr(), 0,
self.cql_grad_scratch.num_bytes(), self.stream.cu_stream(),
self.cql_grad_scratch.len(), self.stream.cu_stream(),
);
}
@@ -7400,11 +7400,11 @@ impl GpuDqnTrainer {
let val_size = (b * na) as i32;
// Zero sensitivity accumulator
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
// Raw cuMemsetD32Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.causal_sensitivity_buf.raw_ptr(), 0,
self.causal_sensitivity_buf.num_bytes(), self.stream.cu_stream(),
self.causal_sensitivity_buf.len(), self.stream.cu_stream(),
);
}
@@ -9152,34 +9152,34 @@ 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 + mse_loss are pinned device-mapped — zero via cuMemsetD8Async on dev_ptr.
// total_loss + mse_loss are pinned device-mapped — zero via cuMemsetD32Async 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::cuMemsetD32Async(
self.total_loss_dev_ptr, 0, 1, self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.mse_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
self.mse_loss_dev_ptr, 0, 1, self.stream.cu_stream(),
);
}
// grad_buf: backward_full uses beta=1.0 GEMM accumulation — zero via ptrs
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.ptrs.grad_buf,
0,
self.total_params * std::mem::size_of::<f32>(),
self.total_params,
self.stream.cu_stream(),
);
}
// d_value/adv_logits: c51_grad + mse_grad kernels write directly (no atomicAdd)
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
// Raw cuMemsetD32Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.d_value_logits_buf.raw_ptr(), 0,
self.d_value_logits_buf.num_bytes(), self.stream.cu_stream(),
self.d_value_logits_buf.len(), self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.d_adv_logits_buf.raw_ptr(), 0,
self.d_adv_logits_buf.num_bytes(), self.stream.cu_stream(),
self.d_adv_logits_buf.len(), self.stream.cu_stream(),
);
}
@@ -9209,15 +9209,15 @@ impl GpuDqnTrainer {
self.launch_curiosity_inference()?;
// MSE path → scratch buffers (REQUIRED: mse_grad_kernel uses atomicAdd)
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
// Raw cuMemsetD32Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.d_value_logits_mse.raw_ptr(), 0,
self.d_value_logits_mse.num_bytes(), self.stream.cu_stream(),
self.d_value_logits_mse.len(), self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.d_adv_logits_mse.raw_ptr(), 0,
self.d_adv_logits_mse.num_bytes(), self.stream.cu_stream(),
self.d_adv_logits_mse.len(), self.stream.cu_stream(),
);
}
self.launch_mse_loss()?;
@@ -9226,8 +9226,8 @@ impl GpuDqnTrainer {
// C51 path → main buffers (already zeroed above)
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.q_divergence_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
self.q_divergence_dev_ptr, 0, 1, self.stream.cu_stream(),
);
}
self.fill_gamma_buf()?;
@@ -9344,20 +9344,20 @@ 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 — pinned device-mapped, use cuMemsetD8Async on dev_ptr.
// Zero accumulators — pinned device-mapped, use cuMemsetD32Async 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::cuMemsetD32Async(
self.total_loss_dev_ptr, 0, 1, self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.mse_loss_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
self.mse_loss_dev_ptr, 0, 1, self.stream.cu_stream(),
);
}
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.ptrs.grad_buf,
0,
self.total_params * std::mem::size_of::<f32>(),
self.total_params,
self.stream.cu_stream(),
);
}
@@ -9379,8 +9379,8 @@ impl GpuDqnTrainer {
// C51 path
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.q_divergence_dev_ptr, 0, std::mem::size_of::<f32>(), self.stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
self.q_divergence_dev_ptr, 0, 1, self.stream.cu_stream(),
);
}
self.fill_gamma_buf()?;

View File

@@ -682,8 +682,8 @@ impl GpuIqlTrainer {
// Zero grad_buf — raw cuMemsetD8Async, no cudarc overhead during graph capture
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.grad_buf.raw_ptr(), 0, self.grad_buf.num_bytes(),
cudarc::driver::sys::cuMemsetD32Async(
self.grad_buf.raw_ptr(), 0, self.grad_buf.len(),
self.stream.cu_stream(),
);
}

View File

@@ -1075,9 +1075,9 @@ impl GpuIqnHead {
// ── Step 6: Quantile Huber loss + dq gradients ──────────────────
// Zero d_branch_logits_buf — raw cuMemsetD8Async, no cudarc overhead during graph capture
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
cudarc::driver::sys::cuMemsetD32Async(
self.d_branch_logits_buf.raw_ptr(), 0,
self.d_branch_logits_buf.num_bytes(),
self.d_branch_logits_buf.len(),
effective_stream.cu_stream(),
);
}
@@ -1113,8 +1113,8 @@ impl GpuIqnHead {
// 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>(), effective_stream.cu_stream(),
cudarc::driver::sys::cuMemsetD32Async(
self.total_loss_dev_ptr, 0, 1, effective_stream.cu_stream(),
);
let loss_ptr = self.total_loss_dev_ptr;
effective_stream
@@ -1184,12 +1184,12 @@ impl GpuIqnHead {
// Zero grad_buf + d_combined_buf — raw cuMemsetD8Async, no cudarc overhead
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.grad_buf.raw_ptr(), 0, self.grad_buf.num_bytes(),
cudarc::driver::sys::cuMemsetD32Async(
self.grad_buf.raw_ptr(), 0, self.grad_buf.len(),
effective_stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.d_combined_buf.raw_ptr(), 0, self.d_combined_buf.num_bytes(),
cudarc::driver::sys::cuMemsetD32Async(
self.d_combined_buf.raw_ptr(), 0, self.d_combined_buf.len(),
effective_stream.cu_stream(),
);
}