diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 348551775..fb34d0fe5 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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); diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 47c3c95f0..5f4c16887 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -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(()) };