diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 1dfe6e962..0b3ff8e49 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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, 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.