perf: add sub-graph phase events inside mega-graph capture

Split fwd_bwd timing into 5 sub-phases: spectral, forward, loss,
backward, aux. Events recorded during graph capture replay with the
graph on every step.

Also extract submit_loss_and_grad_ops() from submit_forward_ops_main()
and make launch_cublas_forward/backward pub(crate) for sub-graph timing.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-11 02:31:03 +02:00
parent a04fcb31f5
commit ef5be1b7a3
2 changed files with 127 additions and 6 deletions

View File

@@ -4591,6 +4591,77 @@ impl GpuDqnTrainer {
Ok(())
}
/// Submit loss computation + gradient ops (everything between forward and backward).
/// Extracted from submit_forward_ops_main for sub-graph timing.
pub(crate) fn submit_loss_and_grad_ops(&mut self) -> Result<(), MLError> {
// Zero accumulators
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}")))?;
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}")))?;
self.launch_curiosity_inference()?;
// MSE path
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
self.launch_c51_loss()?;
self.launch_c51_mixup()?;
self.launch_c51_grad()?;
// Blend: main = α * C51 + (1-α) * MSE
{
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 b3 = self.config.branch_3_size;
let n_val = (b * pad32(na)) as i32;
let n_adv = (b * (b0 + b1 + b2 + b3) * na + 32 * 4) as i32;
let scale_mse = 1.0 - alpha;
let val_blocks = ((n_val as usize + 255) / 256) as u32;
let adv_blocks = ((n_adv as usize + 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 };
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 {
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}")))?;
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}")))?;
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}")))?;
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}")))?;
}
}
Ok(())
}
/// Submit Pass 3 (Double DQN online forward on next_states) for
/// the double_dqn_stream graph. Uses its own CublasForward handle.
///
@@ -4827,7 +4898,7 @@ impl GpuDqnTrainer {
/// uses tensor cores and efficiently tiles across the full SM.
///
/// Returns `Ok(false)` if cuBLAS is not initialized (caller falls back to BF16 kernel).
fn launch_cublas_forward(&self) -> Result<(), MLError> {
pub(crate) fn launch_cublas_forward(&self) -> Result<(), MLError> {
let cublas = &self.cublas_forward;
// Compute 20 raw device pointers from CachedPtrs (no device_ptr calls — graph-safe).
@@ -5471,7 +5542,7 @@ impl GpuDqnTrainer {
/// `d_adv_logits_buf`) directly. The entire backward
/// chain stays in f32 to preserve gradient precision at large batch sizes.
/// Accumulates weight gradients into `grad_buf`.
fn launch_cublas_backward(&self) -> Result<(), MLError> {
pub(crate) fn launch_cublas_backward(&self) -> Result<(), MLError> {
let bw = &self.cublas_backward;
let param_sizes = compute_param_sizes(&self.config);

View File

@@ -115,6 +115,8 @@ pub(crate) const ENS_GRAD_BUDGET: f32 = 0.05;
/// across in-place updates, so the CUDA Graph stays valid.
/// Per-phase CUDA event pairs for GPU-side timing (always-on, zero overhead).
/// Events are hardware timestamp queries — no synchronization, no pipeline stalls.
/// Sub-graph events (spectral/forward/loss/backward/aux) are recorded INSIDE the
/// mega-graph capture — they replay with the graph and give per-phase kernel timing.
struct PhaseEvents {
upload_start: cuda_sys::CUevent,
upload_end: cuda_sys::CUevent,
@@ -124,6 +126,11 @@ struct PhaseEvents {
adam_end: cuda_sys::CUevent,
per_update_start: cuda_sys::CUevent,
per_update_end: cuda_sys::CUevent,
// Sub-graph events (inside mega-graph, replayed on every step)
spectral_end: cuda_sys::CUevent,
forward_end: cuda_sys::CUevent,
loss_end: cuda_sys::CUevent,
backward_end: cuda_sys::CUevent,
}
impl PhaseEvents {
@@ -148,6 +155,10 @@ impl Drop for PhaseEvents {
let _ = cuda_sys::cuEventDestroy_v2(self.adam_end);
let _ = cuda_sys::cuEventDestroy_v2(self.per_update_start);
let _ = cuda_sys::cuEventDestroy_v2(self.per_update_end);
let _ = cuda_sys::cuEventDestroy_v2(self.spectral_end);
let _ = cuda_sys::cuEventDestroy_v2(self.forward_end);
let _ = cuda_sys::cuEventDestroy_v2(self.loss_end);
let _ = cuda_sys::cuEventDestroy_v2(self.backward_end);
}
}
}
@@ -667,6 +678,10 @@ impl FusedTrainingCtx {
adam_end: create_event()?,
per_update_start: create_event()?,
per_update_end: create_event()?,
spectral_end: create_event()?,
forward_end: create_event()?,
loss_end: create_event()?,
backward_end: create_event()?,
})
})() {
Ok(pe) => Some(pe),
@@ -1024,10 +1039,21 @@ impl FusedTrainingCtx {
let per_update = elapsed(pe.per_update_start, pe.per_update_end);
let total = upload + fwd_bwd + adam + per_update;
// Sub-graph breakdown (inside mega-graph replay)
let spectral = elapsed(pe.fwd_bwd_start, pe.spectral_end);
let forward = elapsed(pe.spectral_end, pe.forward_end);
let loss = elapsed(pe.forward_end, pe.loss_end);
let backward = elapsed(pe.loss_end, pe.backward_end);
let aux = elapsed(pe.backward_end, pe.fwd_bwd_end);
tracing::info!(
"GPU phase timing ({} batches, last batch): upload={:.1}ms fwd_bwd={:.1}ms adam={:.1}ms per_update={:.1}ms total={:.1}ms",
self.phase_batch_count, upload, fwd_bwd, adam, per_update, total,
);
tracing::info!(
" fwd_bwd breakdown: spectral={:.1}ms forward={:.1}ms loss={:.1}ms backward={:.1}ms aux={:.1}ms",
spectral, forward, loss, backward, aux,
);
self.phase_batch_count = 0;
}
@@ -1384,15 +1410,39 @@ impl FusedTrainingCtx {
anyhow::bail!("cuStreamBeginCapture_v2 (mega) failed: {begin_result:?}");
}
// Record sub-phase events inside the graph capture — they replay every step.
let cu_stream = self.stream.cu_stream();
// Phase 1: Spectral norm
let spectral_result = self.trainer.apply_spectral_norm(
&self.online_dueling, &self.online_branching,
).map_err(|e| anyhow::anyhow!("mega capture spectral: {e}"));
// Phase 2: Forward + backward (train_step_gpu submits forward ops)
let forward_result = if spectral_result.is_ok() {
self.trainer.submit_forward_ops_main()
.map_err(|e| anyhow::anyhow!("mega capture forward: {e}"))
if let Some(ref pe) = self.phase_events {
PhaseEvents::record(pe.spectral_end, cu_stream);
}
// Phase 2: Forward + loss + backward
let forward_result: Result<()> = if spectral_result.is_ok() {
// Forward GEMMs
self.trainer.launch_cublas_forward()
.map_err(|e| anyhow::anyhow!("mega capture fwd: {e}"))?;
if let Some(ref pe) = self.phase_events {
PhaseEvents::record(pe.forward_end, cu_stream);
}
// Loss + gradient computation
self.trainer.submit_loss_and_grad_ops()
.map_err(|e| anyhow::anyhow!("mega capture loss: {e}"))?;
if let Some(ref pe) = self.phase_events {
PhaseEvents::record(pe.loss_end, cu_stream);
}
// Backward GEMMs
self.trainer.launch_cublas_backward()
.map_err(|e| anyhow::anyhow!("mega capture bwd: {e}"))?;
if let Some(ref pe) = self.phase_events {
PhaseEvents::record(pe.backward_end, cu_stream);
}
Ok(())
} else {
Ok(())
};