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