From 34c8cd1498f9a6f2ea9692b79a6d7366616af627 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 6 Apr 2026 21:20:27 +0200 Subject: [PATCH] fix: convert ALL grad_norm call sites to two-phase reduction MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Previous commit missed 2 IQN/ensemble trunk grad_norm sites that still passed a single-float norm buffer to the new block_sums kernel — buffer overflow on GPU. Now all 4 call sites (main, CQL, IQN trunk, ensemble trunk) use the pre-allocated grad_norm_partials buffer with two-phase reduction. Zero atomicAdd, zero memset across the entire training step. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 63 +++++++++++++------ 1 file changed, 44 insertions(+), 19 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index fd0f614f3..3f05b96cc 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -1196,30 +1196,43 @@ impl GpuDqnTrainer { } // ── 7. Clipped SAXPY: grad_buf[trunk] += iqn_lambda * clip(scratch) ── - // Compute IQN trunk gradient norm, then add with per-component clipping. - // Prevents IQN from overwhelming C51's gradient in grad_buf. + // Two-phase IQN trunk grad norm — no atomicAdd, no memset { - // Zero norm accumulator - self.stream.memset_zeros(&mut self.iqn_trunk_grad_norm) - .map_err(|e| MLError::ModelError(format!("zero iqn_trunk_grad_norm: {e}")))?; - - // Compute IQN trunk gradient norm (sum of squares) let scratch_ptr = self.ptrs.iqn_trunk_m; + let partials_ptr = self.grad_norm_partials.raw_ptr(); let norm_ptr = self.ptrs.iqn_trunk_grad_norm; + let bf16_ptr = self.grad_norm_buf.raw_ptr(); let n_i32 = trunk_grad_total as i32; let blocks = ((trunk_grad_total + 255) / 256) as u32; + let nb = blocks as i32; + // Phase 1: per-block partials unsafe { self.stream .launch_builder(&self.grad_norm_kernel) .arg(&scratch_ptr) - .arg(&norm_ptr) + .arg(&partials_ptr) .arg(&n_i32) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), - shared_mem_bytes: 256, // 8 warps * 2 stride * 4 bytes + shared_mem_bytes: 0, }) - .map_err(|e| MLError::ModelError(format!("IQN trunk grad_norm: {e}")))?; + .map_err(|e| MLError::ModelError(format!("IQN trunk grad_norm phase1: {e}")))?; + } + // Phase 2: reduce → iqn_trunk_grad_norm + unsafe { + self.stream + .launch_builder(&self.grad_norm_finalize_kernel) + .arg(&partials_ptr) + .arg(&norm_ptr) + .arg(&bf16_ptr) + .arg(&nb) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("IQN trunk grad_norm phase2: {e}")))?; } // Clipped SAXPY: grad_buf += iqn_lambda * clip(scratch, iqn_budget) @@ -1443,29 +1456,41 @@ impl GpuDqnTrainer { } // ── 9. Clipped SAXPY: grad_buf[trunk] += scale * clip(scratch) ──── - // Per-component clipping prevents ensemble diversity from overwhelming - // the primary C51 gradient -- same pattern as IQN trunk gradient. + // Two-phase ensemble trunk grad norm — no atomicAdd, no memset { - // Compute ensemble trunk gradient norm - self.stream.memset_zeros(&mut self.iqn_trunk_grad_norm) - .map_err(|e| MLError::ModelError(format!("zero ens_trunk_grad_norm: {e}")))?; - let scratch_ptr = self.ptrs.iqn_trunk_m; + let partials_ptr = self.grad_norm_partials.raw_ptr(); let norm_ptr = self.ptrs.iqn_trunk_grad_norm; + let bf16_ptr = self.grad_norm_buf.raw_ptr(); let n_i32 = trunk_grad_total as i32; let blocks = ((trunk_grad_total + 255) / 256) as u32; + let nb = blocks as i32; unsafe { self.stream .launch_builder(&self.grad_norm_kernel) .arg(&scratch_ptr) - .arg(&norm_ptr) + .arg(&partials_ptr) .arg(&n_i32) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), - shared_mem_bytes: 256, + shared_mem_bytes: 0, }) - .map_err(|e| MLError::ModelError(format!("ens trunk grad_norm: {e}")))?; + .map_err(|e| MLError::ModelError(format!("ens trunk grad_norm phase1: {e}")))?; + } + unsafe { + self.stream + .launch_builder(&self.grad_norm_finalize_kernel) + .arg(&partials_ptr) + .arg(&norm_ptr) + .arg(&bf16_ptr) + .arg(&nb) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("ens trunk grad_norm phase2: {e}")))?; } // Clipped SAXPY