fix: complete CUfunction isolation — zero cross-child sharing + raw memset + pinned memory

Three changes to eliminate the 3100ms replay regression on Hopper:

1. CUfunction isolation: saxpy_f32_kernel was shared across 3 child graphs
   (forward, aux, adam_grad). Added saxpy_f32_adam_grad and saxpy_f32_aux from
   separate CUmodules. Also isolated grad_norm_standalone for post_aux_child
   (was shared with forward_child's d_logits clipping path).

2. Raw cuMemsetD8Async: replaced all cudarc memset_zeros in graph-captured
   functions (submit_forward_ops_main, apply_cql_gradient,
   run_causal_intervention_unconditional) with raw cuMemsetD8Async which is
   properly captured by CUDA Graph. 6 call sites fixed.

3. DtoD memcpy audit: all memcpy_dtod_async calls verified — large buffer
   copies (grad snapshot 2.6MB, multi-horizon blend) are correct for DtoD;
   no scalar copies found that should use pinned memory.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-18 23:24:52 +02:00
parent 79fce72afb
commit 19b6b5af47

View File

@@ -621,6 +621,9 @@ pub struct GpuDqnTrainer {
/// CUDA doesn't allow the same CUfunction to be launched both inside a captured
/// graph and outside on the same stream.
grad_norm_standalone: CudaFunction,
/// Separate grad_norm instance for post_aux_child — cannot share CUfunction
/// with forward_child's grad_norm_standalone on Hopper.
grad_norm_standalone_post_aux: CudaFunction,
adam_update_kernel: CudaFunction,
/// Separate instances for post_aux_child — each from its own CUmodule.
/// CUfunction sharing across child graphs corrupts kernel state on Hopper,
@@ -630,6 +633,11 @@ pub struct GpuDqnTrainer {
ema_kernel: CudaFunction,
saxpy_kernel: CudaFunction,
saxpy_f32_kernel: CudaFunction,
/// Separate saxpy_f32 instance for adam_grad_child — cannot share CUfunction
/// across child graphs (causes 3100ms replay on Hopper).
saxpy_f32_adam_grad: CudaFunction,
/// Separate saxpy_f32 instance for aux_child — same Hopper CUfunction isolation.
saxpy_f32_aux: CudaFunction,
scale_f32_kernel: CudaFunction,
zero_kernel: CudaFunction,
regime_scale_kernel: CudaFunction,
@@ -3004,10 +3012,12 @@ impl GpuDqnTrainer {
let lr = base_lr * self.new_component_lr_scale();
// SGD: params += (-lr) * grad
// Uses adam_grad_child-specific handle — saxpy_f32_kernel is captured in
// forward_child and aux_child; sharing across children corrupts on Hopper.
let mamba2_params_ptr = self.mamba2_params.raw_ptr();
let mamba2_grad_ptr = self.mamba2_grad.raw_ptr();
unsafe {
self.stream.launch_builder(&self.saxpy_f32_kernel)
self.stream.launch_builder(&self.saxpy_f32_adam_grad)
.arg(&mamba2_params_ptr)
.arg(&mamba2_grad_ptr)
.arg(&(-lr))
@@ -3228,6 +3238,7 @@ impl GpuDqnTrainer {
// ── 7. Plain SAXPY: grad_buf[trunk] += iqn_lambda * scratch ──
// Adaptive lambda: scales with IQN loss readiness (0 when uncertain, 1 when converged).
// Uses aux_child-specific handle — saxpy_f32_kernel is captured in forward_child.
{
let scratch_ptr = self.ptrs.iqn_trunk_m;
let grad_ptr = self.ptrs.grad_buf;
@@ -3236,7 +3247,7 @@ impl GpuDqnTrainer {
let blocks = ((trunk_grad_total + 255) / 256) as u32;
unsafe {
self.stream
.launch_builder(&self.saxpy_f32_kernel)
.launch_builder(&self.saxpy_f32_aux)
.arg(&grad_ptr)
.arg(&scratch_ptr)
.arg(&scale)
@@ -3439,6 +3450,7 @@ impl GpuDqnTrainer {
// ── 9. Plain SAXPY: grad_buf[trunk] += scale * scratch ────
// No per-component clip — single global clip in Adam handles safety.
// Uses aux_child-specific handle — saxpy_f32_kernel is captured in forward_child.
{
let scratch_ptr = self.ptrs.iqn_trunk_m;
let grad_ptr = self.ptrs.grad_buf;
@@ -3446,7 +3458,7 @@ impl GpuDqnTrainer {
let blocks = ((trunk_grad_total + 255) / 256) as u32;
unsafe {
self.stream
.launch_builder(&self.saxpy_f32_kernel)
.launch_builder(&self.saxpy_f32_aux)
.arg(&grad_ptr)
.arg(&scratch_ptr)
.arg(&scale)
@@ -3707,8 +3719,13 @@ 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)
self.stream.memset_zeros(&mut self.cql_grad_scratch)
.map_err(|e| MLError::ModelError(format!("zero cql_grad_scratch: {e}")))?;
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.cql_grad_scratch.raw_ptr(), 0,
self.cql_grad_scratch.num_bytes(), self.stream.cu_stream(),
);
}
// GLU saved activation pointers
let glu_gate_pre_ptrs = [
@@ -3796,9 +3813,10 @@ impl GpuDqnTrainer {
let alpha = 1.0_f32;
// Plain SAXPY: grad_buf += 1.0 * cql_scratch
// Uses aux_child-specific handle — saxpy_f32_kernel is captured in forward_child.
unsafe {
self.stream
.launch_builder(&self.saxpy_f32_kernel)
.launch_builder(&self.saxpy_f32_aux)
.arg(&grad_ptr)
.arg(&scratch_ptr)
.arg(&alpha)
@@ -4088,7 +4106,7 @@ impl GpuDqnTrainer {
// per array. Stack is set once in DQNTrainer::new() (64KB for all kernels).
// ── Compile utility kernels (grad_norm, adam_update, etc.) ─
let (grad_norm_kernel, grad_norm_finalize_kernel, grad_norm_finalize_adam, grad_norm_finalize_post_aux, adam_update_kernel, adam_update_post_aux, scale_f32_post_aux, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, scale_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, vaccine_dot_finalize, causal_intervene_kernel_fn, causal_reduce_kernel_fn, causal_mean_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel, nan_check_f32_kernel, nan_check_f32_kernel_b, popart_normalize_kernel) =
let (grad_norm_kernel, grad_norm_finalize_kernel, grad_norm_finalize_adam, grad_norm_finalize_post_aux, adam_update_kernel, adam_update_post_aux, scale_f32_post_aux, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, scale_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, vaccine_dot_finalize, causal_intervene_kernel_fn, causal_reduce_kernel_fn, causal_mean_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel, nan_check_f32_kernel, nan_check_f32_kernel_b, popart_normalize_kernel, saxpy_f32_adam_grad, saxpy_f32_aux, grad_norm_standalone_post_aux) =
compile_training_kernels(&stream, &config)?;
// Separate grad_norm instance for non-graph launches (outside CUDA graph).
@@ -5689,12 +5707,15 @@ impl GpuDqnTrainer {
grad_norm_finalize_adam,
grad_norm_finalize_post_aux,
grad_norm_standalone,
grad_norm_standalone_post_aux,
adam_update_kernel,
adam_update_post_aux,
scale_f32_post_aux,
ema_kernel,
saxpy_kernel,
saxpy_f32_kernel,
saxpy_f32_adam_grad,
saxpy_f32_aux,
scale_f32_kernel,
zero_kernel,
regime_scale_kernel,
@@ -6754,8 +6775,13 @@ impl GpuDqnTrainer {
let val_size = (b * na) as i32;
// Zero sensitivity accumulator
self.stream.memset_zeros(&mut self.causal_sensitivity_buf)
.map_err(|e| MLError::ModelError(format!("zero causal_sens: {e}")))?;
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.causal_sensitivity_buf.raw_ptr(), 0,
self.causal_sensitivity_buf.num_bytes(), self.stream.cu_stream(),
);
}
let param_sizes = compute_param_sizes(&self.config);
let on_w_ptrs = f32_weight_ptrs_from_base(self.ptrs.params_ptr, &param_sizes);
@@ -7823,11 +7849,13 @@ impl GpuDqnTrainer {
unsafe { *self.denoise_t_pinned = step_val; }
// Phase 1: grad_norm on denoise_grad
// Uses post_aux_child-specific handle — grad_norm_standalone is captured in
// forward_child (d_logits clipping). Sharing across children corrupts on Hopper.
let grad_ptr = self.denoise_grad.raw_ptr();
let partials_ptr = self.denoise_norm_partials.raw_ptr();
unsafe {
self.stream
.launch_builder(&self.grad_norm_standalone)
.launch_builder(&self.grad_norm_standalone_post_aux)
.arg(&grad_ptr)
.arg(&partials_ptr)
.arg(&n)
@@ -7959,12 +7987,17 @@ impl GpuDqnTrainer {
);
}
// d_value/adv_logits: c51_grad + mse_grad kernels write directly (no atomicAdd)
self.stream
.memset_zeros(&mut self.d_value_logits_buf)
.map_err(|e| MLError::ModelError(format!("zero d_value_logits: {e}")))?;
self.stream
.memset_zeros(&mut self.d_adv_logits_buf)
.map_err(|e| MLError::ModelError(format!("zero d_adv_logits: {e}")))?;
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.d_value_logits_buf.raw_ptr(), 0,
self.d_value_logits_buf.num_bytes(), self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.d_adv_logits_buf.raw_ptr(), 0,
self.d_adv_logits_buf.num_bytes(), self.stream.cu_stream(),
);
}
// ── 1. Forward (cuBLAS SGEMM — Pass 1 + 2, no Pass 3) ──────────
self.launch_cublas_forward()?;
@@ -7992,10 +8025,17 @@ impl GpuDqnTrainer {
self.launch_curiosity_inference()?;
// MSE path → scratch buffers (REQUIRED: mse_grad_kernel uses atomicAdd)
self.stream.memset_zeros(&mut self.d_value_logits_mse)
.map_err(|e| MLError::ModelError(format!("zero d_value_mse: {e}")))?;
self.stream.memset_zeros(&mut self.d_adv_logits_mse)
.map_err(|e| MLError::ModelError(format!("zero d_adv_mse: {e}")))?;
// Raw cuMemsetD8Async — cudarc memset_zeros is NOT captured by CUDA Graph.
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.d_value_logits_mse.raw_ptr(), 0,
self.d_value_logits_mse.num_bytes(), self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.d_adv_logits_mse.raw_ptr(), 0,
self.d_adv_logits_mse.num_bytes(), self.stream.cu_stream(),
);
}
self.launch_mse_loss()?;
self.launch_loss_reduce(self.mse_loss_dev_ptr)?;
self.launch_mse_grad_to_scratch()?;
@@ -10168,11 +10208,13 @@ impl GpuDqnTrainer {
unsafe { *self.sel_t_pinned = step_val; }
// Phase 1: grad_norm_kernel on sel_grad → sel_norm_partials [1 block].
// Uses post_aux_child-specific handle — grad_norm_standalone is captured in
// forward_child (d_logits clipping). Sharing across children corrupts on Hopper.
let grad_ptr = self.sel_grad.raw_ptr();
let partials_ptr = self.sel_norm_partials.raw_ptr();
unsafe {
self.stream
.launch_builder(&self.grad_norm_standalone)
.launch_builder(&self.grad_norm_standalone_post_aux)
.arg(&grad_ptr)
.arg(&partials_ptr)
.arg(&sel_n)
@@ -10260,7 +10302,7 @@ impl GpuDqnTrainer {
fn compile_training_kernels(
stream: &Arc<CudaStream>,
config: &GpuDqnTrainConfig,
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
info!(
state_dim = config.state_dim,
total_params = compute_total_params(config),
@@ -10292,6 +10334,23 @@ fn compile_training_kernels(
.map_err(|e| MLError::ModelError(format!("post_aux adam_update load: {e}")))?;
let scale_f32_post_aux = post_aux_module.load_function("dqn_scale_f32_kernel")
.map_err(|e| MLError::ModelError(format!("post_aux scale_f32 load: {e}")))?;
// Separate grad_norm instance for post_aux_child — forward_child already uses
// grad_norm_standalone for d_logits clipping. Sharing across children corrupts
// CUfunction state on Hopper → 3100ms replay.
let grad_norm_standalone_post_aux = post_aux_module.load_function("dqn_grad_norm_kernel")
.map_err(|e| MLError::ModelError(format!("post_aux grad_norm_standalone load: {e}")))?;
// Separate CUmodule for adam_grad_child — saxpy_f32 is used in forward_child and
// aux_child, so adam_grad_child needs its own isolated handle.
let adam_grad_module = context.load_cubin(DQN_UTILITY_CUBIN.to_vec())
.map_err(|e| MLError::ModelError(format!("adam_grad cubin load: {e}")))?;
let saxpy_f32_adam_grad = adam_grad_module.load_function("dqn_saxpy_f32_kernel")
.map_err(|e| MLError::ModelError(format!("adam_grad saxpy_f32 load: {e}")))?;
// Separate CUmodule for aux_child — saxpy_f32 is used in forward_child and
// adam_grad_child, so aux_child needs its own isolated handle.
let aux_module = context.load_cubin(DQN_UTILITY_CUBIN.to_vec())
.map_err(|e| MLError::ModelError(format!("aux cubin load: {e}")))?;
let saxpy_f32_aux = aux_module.load_function("dqn_saxpy_f32_kernel")
.map_err(|e| MLError::ModelError(format!("aux saxpy_f32 load: {e}")))?;
let adam_update = module.load_function("dqn_adam_update_kernel")
.map_err(|e| MLError::ModelError(format!("dqn_adam_update_kernel load: {e}")))?;
// f32_to_bf16_kernel / bf16_to_f32_kernel DELETED — pure f32 pipeline.
@@ -10352,8 +10411,8 @@ fn compile_training_kernels(
let popart_normalize = module.load_function("popart_normalize_kernel")
.map_err(|e| MLError::ModelError(format!("popart_normalize_kernel load: {e}")))?;
info!("GpuDqnTrainer: 30 utility kernels loaded from precompiled cubin");
Ok((grad_norm, grad_norm_finalize, grad_norm_finalize_adam, grad_norm_finalize_post_aux, adam_update, adam_update_post_aux, scale_f32_post_aux, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clip_grad, pad_states, saxpy_f32, scale_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, vaccine_dot_finalize, causal_intervene, causal_reduce, causal_mean_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace, nan_check_f32, nan_check_f32_b, popart_normalize))
info!("GpuDqnTrainer: 33 utility kernels loaded from precompiled cubin (5 CUmodules)");
Ok((grad_norm, grad_norm_finalize, grad_norm_finalize_adam, grad_norm_finalize_post_aux, adam_update, adam_update_post_aux, scale_f32_post_aux, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clip_grad, pad_states, saxpy_f32, scale_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, vaccine_dot_finalize, causal_intervene, causal_reduce, causal_mean_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace, nan_check_f32, nan_check_f32_b, popart_normalize, saxpy_f32_adam_grad, saxpy_f32_aux, grad_norm_standalone_post_aux))
}
/// Load the standalone Polyak EMA kernel from precompiled cubin.