From 13dd2e77bfdcd6d1331f5a107f01178fc7344d7a Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Thu, 2 Apr 2026 10:06:56 +0200 Subject: [PATCH] =?UTF-8?q?perf:=20graph=20all=20remaining=20ops=20?= =?UTF-8?q?=E2=80=94=20IQN=20full=20pipeline=20+=20regime=5Fscale=20in=20g?= =?UTF-8?q?raph=5Fadam?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit IQN graph now captures the FULL pipeline in one graph: - decode_actions + fwd/loss + backward + grad_norm + Adam - trunk gradient (cuBLAS backward into shared weights) - target EMA (tau from GPU-resident tau_buf, async HtoD before replay) - IQN→PER loss DtoD copy IQN EMA kernel changed: float tau → const float* tau_buf (device read). tau_buf added to GpuIqnHead with async cuMemcpyHtoDAsync per step. This was the last scalar parameter preventing full graph capture. regime_scale_td_errors moved into graph_adam submit sequence. Runs after Adam unflatten, before PER priority update. Per-step: 7 graph replays + ~9 ungraphed ops Ungraphed ops (genuinely can't be graphed — batch ptrs change): - upload_batch_gpu: 6 DtoD + 2 pad_states (batch-specific pointers) - HER relabel: 1-2 kernels (donor from batch next_states) - PER priority update: 1 kernel (indices from batch) Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 3 + crates/ml/src/cuda_pipeline/gpu_iqn_head.rs | 17 +++- .../src/cuda_pipeline/iqn_dual_head_kernel.cu | 10 +-- crates/ml/src/trainers/dqn/fused_training.rs | 83 +++++++++++-------- 4 files changed, 71 insertions(+), 42 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 171e03576..f3a1fe06c 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -4348,6 +4348,9 @@ impl GpuDqnTrainer { // ── 7. Unflatten: params_bf16 → individual bf16 weight tensors ─ self.unflatten_online_weights(online_d, online_b)?; + // ── 8. Regime-adaptive PER scaling (element-wise, all pre-allocated) ─ + self.regime_scale_td_errors()?; + Ok(()) } diff --git a/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs b/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs index d3516ade2..b6ba0a0c1 100644 --- a/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs +++ b/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs @@ -189,6 +189,7 @@ pub struct GpuIqnHead { // ── Training state ─────────────────────────────────────────────── adam_step: i32, t_buf: cudarc::driver::CudaSlice, + tau_buf: cudarc::driver::CudaSlice, /// Monotonic step counter for Philox PRNG seeding (τ sampling). rng_step: u32, total_params: usize, @@ -279,6 +280,8 @@ impl GpuIqnHead { let t_buf = stream.alloc_zeros::(1) .map_err(|e| MLError::ModelError(format!("iqn_t_buf alloc: {e}")))?; + let tau_buf = stream.alloc_zeros::(1) + .map_err(|e| MLError::ModelError(format!("iqn_tau_buf alloc: {e}")))?; Ok(Self { config, stream, @@ -312,6 +315,7 @@ impl GpuIqnHead { total_loss, adam_step: 0, t_buf, + tau_buf, rng_step: 0, total_params, }) @@ -591,6 +595,14 @@ impl GpuIqnHead { /// `target[i] = (1 - tau) * target[i] + tau * online[i]` pub fn target_ema_update(&mut self, tau: f32) -> Result<(), MLError> { let n = self.total_params; + unsafe { + cudarc::driver::sys::cuMemcpyHtoDAsync_v2( + self.tau_buf.raw_ptr(), + (&tau as *const f32).cast(), + std::mem::size_of::(), + self.stream.cu_stream(), + ); + } let blocks = (n + 255) / 256; let config = LaunchConfig { grid_dim: (blocks as u32, 1, 1), @@ -602,6 +614,7 @@ impl GpuIqnHead { let shared_h1_i32 = self.config.shared_h1 as i32; let hidden_dim_i32 = self.config.hidden_dim as i32; let embed_dim_i32 = self.config.embed_dim as i32; + let tau_ptr = self.tau_buf.raw_ptr(); // Safety: target_params and online_params have total_params elements each. unsafe { @@ -609,7 +622,7 @@ impl GpuIqnHead { .launch_builder(&self.ema_kernel) .arg(&mut self.target_params) .arg(&self.online_params) - .arg(&tau) + .arg(&tau_ptr) .arg(&n_i32) .arg(&shared_h1_i32) .arg(&hidden_dim_i32) @@ -639,6 +652,8 @@ impl GpuIqnHead { } /// Raw device pointer to IQN trunk gradient — avoids CudaSlice borrow conflicts. + pub fn tau_buf_ptr(&self) -> u64 { self.tau_buf.raw_ptr() } + pub fn d_h_s2_raw_ptr(&self) -> u64 { self.d_h_s2_buf.raw_ptr() } diff --git a/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu b/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu index 265e3edae..d3b6419e5 100644 --- a/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu +++ b/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu @@ -875,16 +875,16 @@ extern "C" __global__ void iqn_ema_kernel( __nv_bfloat16* __restrict__ target, const __nv_bfloat16* __restrict__ online, - float tau, + const float* __restrict__ tau_buf, int n, - int shared_h1, /* runtime: unused, for consistent interface */ - int hidden_dim, /* runtime: unused, for consistent interface */ - int embed_dim /* runtime: unused, for consistent interface */ + int shared_h1, + int hidden_dim, + int embed_dim ) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < n) { - __nv_bfloat16 bf16_tau = bf16(tau); + __nv_bfloat16 bf16_tau = bf16(tau_buf[0]); target[i] = (bf16_one() - bf16_tau) * target[i] + bf16_tau * online[i]; } } diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index d45e07c73..cdc80f984 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -765,14 +765,31 @@ impl FusedTrainingCtx { // IQN: most of the pipeline is graphed, but trunk gradient + target EMA // use tau (changes per step) and write to online_dueling (shared state). // Graph the main IQN forward+loss+backward+adam, keep trunk grad + EMA ungraphed. + // IQN: full pipeline + trunk gradient + EMA + PER DtoD in one graph. + // tau uploaded async before replay (GPU-resident tau_buf). if self.graph_iqn.is_none() && self.gpu_iqn.is_some() { + // First step: run everything ungraphed if let Some(ref mut iqn) = self.gpu_iqn { let _ = iqn.train_iqn_step_gpu( self.trainer.save_h_s2(), self.trainer.next_states_buf(), &self.target_dueling, self.trainer.actions_buf(), self.trainer.rewards_buf(), self.trainer.dones_buf(), ); + let d_ptr = iqn.d_h_s2_raw_ptr(); + let _ = self.trainer.apply_iqn_trunk_gradient(d_ptr, &mut self.online_dueling); + let _ = iqn.target_ema_update(0.005); // initial tau + // IQN→PER DtoD + let bs = self.trainer.batch_size(); + let n_bytes = bs * std::mem::size_of::(); + unsafe { + let _ = cudarc::driver::result::memcpy_dtod_async( + self.trainer.td_errors_buf().raw_ptr(), + iqn.per_sample_loss().raw_ptr(), + n_bytes, self.stream.cu_stream(), + ); + } } + // Capture the full IQN pipeline as one graph self.stream.synchronize() .map_err(|e| anyhow::anyhow!("sync before iqn capture: {e}"))?; unsafe { self.stream.context().disable_event_tracking(); } @@ -785,49 +802,47 @@ impl FusedTrainingCtx { &self.target_dueling, self.trainer.actions_buf(), self.trainer.rewards_buf(), self.trainer.dones_buf(), ); + let d_ptr = iqn.d_h_s2_raw_ptr(); + let _ = self.trainer.apply_iqn_trunk_gradient(d_ptr, &mut self.online_dueling); + let _ = iqn.target_ema_update(0.005); + let bs = self.trainer.batch_size(); + let n_bytes = bs * std::mem::size_of::(); + unsafe { + let _ = cudarc::driver::result::memcpy_dtod_async( + self.trainer.td_errors_buf().raw_ptr(), + iqn.per_sample_loss().raw_ptr(), + n_bytes, self.stream.cu_stream(), + ); + } } if let Ok(Some(graph)) = self.stream.end_capture( cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, ) { use crate::cuda_pipeline::gpu_dqn_trainer::SendSyncGraph; self.graph_iqn = Some(SendSyncGraph(graph)); - tracing::info!("Captured IQN CUDA graph (decode + fwd/loss + bwd + adam)"); + tracing::info!("Captured IQN CUDA graph (train + trunk grad + EMA + PER DtoD)"); } } unsafe { self.stream.context().enable_event_tracking(); } let _ = self.stream.context().check_err(); } else if let Some(ref graph) = self.graph_iqn { - graph.0.launch().map_err(|e| anyhow::anyhow!("IQN graph replay: {e}"))?; - } - - // IQN trunk gradient + target EMA (ungraphed — tau changes per step) - if let Some(ref mut iqn) = self.gpu_iqn { - let d_h_s2_ptr = iqn.d_h_s2_raw_ptr(); - self.trainer.apply_iqn_trunk_gradient( - d_h_s2_ptr, - &mut self.online_dueling, - ).map_err(|e| anyhow::anyhow!("IQN trunk gradient: {e}"))?; - - let dqn = agent.primary_dqn_mut(); - let tau = compute_cosine_annealed_tau( - dqn.get_training_steps(), - dqn.config.tau, dqn.config.tau_final, dqn.config.tau_anneal_steps, - ); - iqn.target_ema_update(tau as f32) - .map_err(|e| anyhow::anyhow!("IQN EMA update: {e}"))?; - } - - // IQN→PER loss DtoD (ungraphed — td_errors_buf address could shift if PER resizes) - if let Some(ref mut iqn) = self.gpu_iqn { - let bs = self.trainer.batch_size(); - let n_bytes = bs * std::mem::size_of::(); - let src_ptr = iqn.per_sample_loss().raw_ptr(); - let dst_ptr = self.trainer.td_errors_buf().raw_ptr(); - unsafe { - cudarc::driver::result::memcpy_dtod_async( - dst_ptr, src_ptr, n_bytes, self.stream.cu_stream() - ).map_err(|e| anyhow::anyhow!("IQN→PER loss DtoD: {e}"))?; + // Upload tau before replay (IQN EMA reads from tau_buf) + if let Some(ref iqn) = self.gpu_iqn { + let dqn = agent.primary_dqn_mut(); + let tau = compute_cosine_annealed_tau( + dqn.get_training_steps(), + dqn.config.tau, dqn.config.tau_final, dqn.config.tau_anneal_steps, + ); + unsafe { + cudarc::driver::sys::cuMemcpyHtoDAsync_v2( + iqn.tau_buf_ptr(), + (&(tau as f32) as *const f32).cast(), + std::mem::size_of::(), + self.stream.cu_stream(), + ); + } } + graph.0.launch().map_err(|e| anyhow::anyhow!("IQN graph replay: {e}"))?; } // ── Step 5b2: Ensemble multi-head diversity loss ────────────── @@ -880,11 +895,7 @@ impl FusedTrainingCtx { let fused_result = self.trainer.replay_adam_and_readback() .map_err(|e| { eprintln!("!!! ADAM REPLAY FAILED: {e}"); anyhow::anyhow!("graph_adam replay: {e}") })?; - // ── Step 5f: Regime-adaptive PER scaling ────────────────────── - // Kernel reads target ADX/CUSUM from first sample in states_buf. - // Zero CPU readback — fully GPU-native. - self.trainer.regime_scale_td_errors() - .map_err(|e| anyhow::anyhow!("Regime PER scaling: {e}"))?; + // Regime PER scaling now captured in graph_adam — no per-step call needed. // ── Step 6: GPU-native PER priority update ───────────────────── // td_errors stay on GPU (td_errors_buf). Single CUDA kernel scatter-writes