From c1e1cd47f7a23b4f14f0509659bed4e10ec934a3 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 5 Apr 2026 22:04:29 +0200 Subject: [PATCH] fix: dual-stream CUDA graph capture eliminates H100 hang MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Split graph_forward into graph_forward_main (Pass 1+2+loss+backward on main stream) and graph_forward_ddqn (Pass 3 on double_dqn_stream). Event synchronization happens in ungraphed replay_forward() code, not inside captured functions. Each graph uses its own CublasForward handle — no set_stream() during capture. Fixes: H100 training hang with gpu_n_episodes=4096. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 171 ++++++------------ 1 file changed, 56 insertions(+), 115 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 2a5b671fe..3c8736986 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -601,6 +601,8 @@ pub struct GpuDqnTrainer { pub(crate) graph_forward: Option, /// MSE-only forward graph (no C51 — faster, avoids NaN during warmup). pub(crate) graph_forward_mse: Option, + /// Double DQN forward graph (Pass 3 on double_dqn_stream). + pub(crate) graph_forward_ddqn: Option, pub(crate) graph_adam: Option, // ── Consolidated transfer buffers ───────────────────────────── @@ -914,6 +916,7 @@ impl Drop for GpuDqnTrainer { unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } self.graph_forward = None; self.graph_forward_mse = None; + self.graph_forward_ddqn = None; self.graph_adam = None; } } @@ -2174,8 +2177,9 @@ impl GpuDqnTrainer { } drop(_evt_guard); - // Invalidate both CUDA Graphs — weights changed outside captured ops + // Invalidate all CUDA Graphs — weights changed outside captured ops self.graph_forward = None; + self.graph_forward_ddqn = None; self.graph_adam = None; info!(alpha, sigma, total_params = n, "Shrink-and-Perturb applied (GPU-native, zero CPU)"); @@ -2789,6 +2793,7 @@ impl GpuDqnTrainer { attention_initialized: false, graph_forward: None, graph_forward_mse: None, + graph_forward_ddqn: None, graph_adam: None, upload_staging_buf, upload_staging_len, @@ -3289,6 +3294,7 @@ impl GpuDqnTrainer { // Invalidate CUDA graphs (mask changes the training step) self.graph_forward = None; + self.graph_forward_ddqn = None; self.graph_adam = None; Ok(()) @@ -3514,17 +3520,39 @@ impl GpuDqnTrainer { } pub fn replay_forward(&self) -> Result<(), MLError> { - // Route to MSE-only graph during warmup (faster, avoids C51 NaN) - let graph = if self.c51_alpha < 1e-6 { + // 1. Replay main stream graph (Pass 1 + 2 + loss + grad + backward) + let main_graph = if self.c51_alpha < 1e-6 { self.graph_forward_mse.as_ref().or(self.graph_forward.as_ref()) } else { self.graph_forward.as_ref() }; - if let Some(g) = graph { + if let Some(g) = main_graph { g.0.launch().map_err(|e| { MLError::ModelError(format!("CUDA graph_forward replay: {e}")) })?; } + + // 2. Record event after main graph completes + self.pass1_event.record(&self.stream) + .map_err(|e| MLError::ModelError(format!("pass1 event record: {e}")))?; + + // 3. double_dqn_stream waits for main stream + self.double_dqn_stream.wait(&self.pass1_event) + .map_err(|e| MLError::ModelError(format!("ddqn wait pass1: {e}")))?; + + // 4. Replay Pass 3 on double_dqn_stream (concurrent with Pass 2 inside main graph) + if let Some(ref g) = self.graph_forward_ddqn { + g.0.launch().map_err(|e| { + MLError::ModelError(format!("CUDA graph_forward_ddqn replay: {e}")) + })?; + } + + // 5. Main stream waits for Pass 3 to complete + self.pass3_event.record(&self.double_dqn_stream) + .map_err(|e| MLError::ModelError(format!("pass3 event record: {e}")))?; + self.stream.wait(&self.pass3_event) + .map_err(|e| MLError::ModelError(format!("join pass3: {e}")))?; + Ok(()) } @@ -4126,7 +4154,7 @@ impl GpuDqnTrainer { return Err(MLError::ModelError(format!("CUDA graph_forward begin_capture: {e}"))); } - let submit_fwd_result = self.submit_forward_ops(); + let submit_fwd_result = self.submit_forward_ops_main(); let graph_fwd_result = self.stream.end_capture( cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, @@ -4183,6 +4211,24 @@ impl GpuDqnTrainer { MLError::ModelError("graph_forward_mse capture returned None".into()) })?; + // ── Capture graph_forward_ddqn (Pass 3 on double_dqn_stream) ── + let begin_ddqn = self.double_dqn_stream.begin_capture( + cudarc::driver::sys::CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_THREAD_LOCAL, + ); + if let Err(e) = begin_ddqn { + unsafe { self.stream.context().enable_event_tracking(); } + let _ = self.stream.context().check_err(); + return Err(MLError::ModelError(format!("graph_forward_ddqn begin_capture: {e}"))); + } + let submit_ddqn_result = self.submit_forward_ops_ddqn(); + let graph_ddqn_result = self.double_dqn_stream.end_capture( + cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, + ); + submit_ddqn_result?; + let graph_ddqn = graph_ddqn_result + .map_err(|e| MLError::ModelError(format!("graph_forward_ddqn end_capture: {e}")))? + .ok_or_else(|| MLError::ModelError("graph_forward_ddqn capture returned None".into()))?; + // ── Capture graph_adam ────────────────────────────────────────── let begin_result = self.stream.begin_capture( cudarc::driver::sys::CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_THREAD_LOCAL, @@ -4218,12 +4264,14 @@ impl GpuDqnTrainer { // downstream buffers, causing NaN at step 0. info!( - "GpuDqnTrainer: 3 CUDA graphs captured \ + "GpuDqnTrainer: 4 CUDA graphs captured \ (graph_forward: MSE+C51; graph_forward_mse: MSE-only; \ + graph_forward_ddqn: Pass 3 on double_dqn_stream; \ graph_adam: grad_norm + adam + unflatten)" ); self.graph_forward = Some(SendSyncGraph(graph_fwd)); self.graph_forward_mse = Some(SendSyncGraph(graph_mse)); + self.graph_forward_ddqn = Some(SendSyncGraph(graph_ddqn)); self.graph_adam = Some(SendSyncGraph(graph_adam)); self.last_captured_loss_mode = Some(self.loss_mode); Ok(()) @@ -4261,115 +4309,6 @@ impl GpuDqnTrainer { Ok(()) } - pub(crate) fn submit_forward_ops(&mut self) -> Result<(), MLError> { - // ── Zero accumulators (capturable: memset_zeros uses cuMemsetD32Async) ─ - self.stream - .memset_zeros(&mut self.total_loss_buf) - .map_err(|e| MLError::ModelError(format!("zero total_loss: {e}")))?; - self.stream - .memset_zeros(&mut self.mse_loss_buf) - .map_err(|e| MLError::ModelError(format!("zero mse_loss: {e}")))?; - self.stream - .memset_zeros(&mut self.grad_buf) - .map_err(|e| MLError::ModelError(format!("zero grad_buf: {e}")))?; - // Zero C51 gradient output buffers (atomicAdd accumulates) - 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}")))?; - - // ── 1. Forward (cuBLAS SGEMM, 3 passes) ────────────────── - self.launch_cublas_forward()?; - - // ── 2+3. Loss + gradient (blended MSE + C51 via c51_alpha ramp) ─ - // Always run BOTH loss paths. Blend gradients: - // grad = (1-α) * MSE_grad + α * C51_grad - // When α=0 (warmup): pure MSE. When α=1 (converged): pure C51. - // The kernel sequence is FIXED (both always run) → CUDA Graph compatible. - // - // MSE writes to d_value_logits_mse / d_adv_logits_mse (scratch) - // C51 writes to d_value_logits_buf / d_adv_logits_buf (main) - // Then SAXPY blends: main = α * main + (1-α) * scratch - - // ── Curiosity Q-penalty: compute per-sample prediction error ── - // Must run BEFORE both MSE and C51 loss kernels (both read curiosity_error_buf). - // Captured in graph — weight pointers are stable (in-place update by curiosity trainer). - self.launch_curiosity_inference()?; - - // MSE path → scratch buffers - 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}")))?; - self.launch_mse_loss()?; - self.launch_mse_grad_to_scratch()?; - - // C51 path → main buffers (already zeroed above) - // Zero Manifold Mixup inter-block barrier counters [NUM_BRANCHES=3] before C51 launch. - self.stream.memset_zeros(&mut self.mixup_barrier_buf) - .map_err(|e| MLError::ModelError(format!("zero mixup_barrier: {e}")))?; - self.launch_c51_loss()?; - self.launch_c51_grad()?; - - // Blend: main = α * C51 + (1-α) * MSE - // Step 1: main *= α (scale C51 grad) - // Step 2: main += (1-α) * scratch (add MSE grad) - { - let alpha = self.c51_alpha; - let na = self.config.num_atoms; - let b = self.config.batch_size; - let b0 = self.config.branch_0_size; - let b1 = self.config.branch_1_size; - let b2 = self.config.branch_2_size; - let n_val = (b * na) as i32; - let n_adv = (b * (b0 + b1 + b2) * na) as i32; - let scale_mse = 1.0 - alpha; - let val_blocks = ((b * na + 255) / 256) as u32; - let adv_blocks = ((b * (b0 + b1 + b2) * na + 255) / 256) as u32; - let cfg_val = LaunchConfig { grid_dim: (val_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }; - let cfg_adv = LaunchConfig { grid_dim: (adv_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }; - - // Use raw device pointers to avoid &mut/& borrow conflict on same buffer. - let val_ptr = self.d_value_logits_buf.raw_ptr(); - let val_mse_ptr = self.d_value_logits_mse.raw_ptr(); - let adv_ptr = self.d_adv_logits_buf.raw_ptr(); - let adv_mse_ptr = self.d_adv_logits_mse.raw_ptr(); - - unsafe { - // d_value *= α (in-place scale — single pointer, no aliasing) - self.stream.launch_builder(&self.scale_f32_kernel) - .arg(&val_ptr) - .arg(&alpha).arg(&n_val) - .launch(cfg_val).map_err(|e| MLError::ModelError(format!("scale c51 val: {e}")))?; - // d_value += (1-α) * mse_scratch (SAXPY with distinct pointers — safe) - self.stream.launch_builder(&self.saxpy_f32_kernel) - .arg(&val_ptr).arg(&val_mse_ptr) - .arg(&scale_mse).arg(&n_val) - .launch(cfg_val).map_err(|e| MLError::ModelError(format!("blend mse val: {e}")))?; - // d_adv *= α (in-place scale — single pointer, no aliasing) - self.stream.launch_builder(&self.scale_f32_kernel) - .arg(&adv_ptr) - .arg(&alpha).arg(&n_adv) - .launch(cfg_adv).map_err(|e| MLError::ModelError(format!("scale c51 adv: {e}")))?; - // d_adv += (1-α) * mse_scratch (SAXPY with distinct pointers — safe) - self.stream.launch_builder(&self.saxpy_f32_kernel) - .arg(&adv_ptr).arg(&adv_mse_ptr) - .arg(&scale_mse).arg(&n_adv) - .launch(cfg_adv).map_err(|e| MLError::ModelError(format!("blend mse adv: {e}")))?; - } - } - - // ── 3.5. Cast f32 d_logits → bf16 staging for cuBLAS backward ─ - self.cast_d_logits_to_bf16()?; - - // ── 4. Backward (cuBLAS SGEMM, chain rule through layers) ─ - self.launch_cublas_backward()?; - - Ok(()) - } - /// Submit main-stream forward ops for graph capture (Pass 1 + 2 + loss + grad + backward). /// /// Contains everything that runs on `self.stream`: @@ -4558,6 +4497,7 @@ impl GpuDqnTrainer { pub fn invalidate_training_graph(&mut self) { self.graph_forward = None; self.graph_forward_mse = None; + self.graph_forward_ddqn = None; self.graph_adam = None; self.last_captured_loss_mode = None; self.params_initialized = false; @@ -4849,6 +4789,7 @@ impl GpuDqnTrainer { self.curiosity_b2_ptr = b2.raw_ptr(); // Invalidate CUDA Graphs so next capture includes curiosity inference self.graph_forward = None; + self.graph_forward_ddqn = None; self.graph_adam = None; }