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:
@@ -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, ¶m_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.
|
||||
|
||||
Reference in New Issue
Block a user