fix: dual-stream CUDA graph capture eliminates H100 hang

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-05 22:04:29 +02:00
parent f59c615687
commit c1e1cd47f7

View File

@@ -601,6 +601,8 @@ pub struct GpuDqnTrainer {
pub(crate) graph_forward: Option<SendSyncGraph>,
/// MSE-only forward graph (no C51 — faster, avoids NaN during warmup).
pub(crate) graph_forward_mse: Option<SendSyncGraph>,
/// Double DQN forward graph (Pass 3 on double_dqn_stream).
pub(crate) graph_forward_ddqn: Option<SendSyncGraph>,
pub(crate) graph_adam: Option<SendSyncGraph>,
// ── 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;
}