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:
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user