fix: convert ALL grad_norm call sites to two-phase reduction

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-06 21:20:27 +02:00
parent 4fcfb32637
commit 34c8cd1498

View File

@@ -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