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