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:
@@ -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(),
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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()?;
|
||||
|
||||
@@ -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(),
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user