fix: eliminate grad_buf pointer duplication — avg_grad=0 root cause

All code now uses ptrs.grad_buf (u64) as single source of truth.
Vaccine rewritten to use launch_cublas_backward_to(scratch) instead
of std::mem::swap on CudaSlice which caused pointer divergence
between diagnostics and Adam/graphs. Locally verified: avg_grad=6.13.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-13 23:05:43 +02:00
parent 03cac6ffda
commit 4159d201a0
2 changed files with 43 additions and 55 deletions

View File

@@ -2135,9 +2135,7 @@ impl GpuDqnTrainer {
let d_val_norm = self.debug_buffer_norm_f32(self.d_value_logits_buf.raw_ptr(), b * na).unwrap_or(f32::NAN);
let d_adv_norm = self.debug_buffer_norm_f32(self.d_adv_logits_buf.raw_ptr(), b * tba).unwrap_or(f32::NAN);
let grad_norm = self.debug_buffer_norm_f32(self.grad_buf.raw_ptr(), self.total_params).unwrap_or(f32::NAN);
// h_s2 is f32 — could add direct NaN check, but the f32 buffers tell us enough.
// Skip for now — the f32 buffers tell us enough.
let grad_norm = self.debug_buffer_norm_f32(self.ptrs.grad_buf, self.total_params).unwrap_or(f32::NAN);
tracing::warn!(
step,
@@ -2166,7 +2164,7 @@ impl GpuDqnTrainer {
let d_adv_norm = self.debug_buffer_norm_f32(self.d_adv_logits_buf.raw_ptr(), b * tba).unwrap_or(f32::NAN);
let mse_val_norm = self.debug_buffer_norm_f32(self.d_value_logits_mse.raw_ptr(), b * na).unwrap_or(f32::NAN);
let mse_adv_norm = self.debug_buffer_norm_f32(self.d_adv_logits_mse.raw_ptr(), b * tba).unwrap_or(f32::NAN);
let grad_norm = self.debug_buffer_norm_f32(self.grad_buf.raw_ptr(), self.total_params).unwrap_or(f32::NAN);
let grad_norm = self.debug_buffer_norm_f32(self.ptrs.grad_buf, self.total_params).unwrap_or(f32::NAN);
tracing::warn!(
remaining = self.diag_recapture_remaining,
@@ -3297,14 +3295,6 @@ impl GpuDqnTrainer {
/// Replay graph_adam + readback scalars. Call AFTER injecting auxiliary
/// gradients into grad_buf (IQN, attention, ensemble).
/// Restore grad_buf after a failed vaccine step by swapping back.
/// Called when apply_gradient_vaccine fails mid-step (grad_buf and
/// vaccine_grad_save are in swapped state). Zero-copy — just swaps
/// the CudaSlice structs to restore the original pointers.
pub fn restore_grad_from_vaccine_save(&mut self) -> Result<(), MLError> {
std::mem::swap(&mut self.grad_buf, &mut self.vaccine_grad_save);
Ok(())
}
/// Compute gradient norm OUTSIDE the CUDA graph.
/// Two-phase reduction: phase 1 writes per-block partials (no atomicAdd),
@@ -3771,39 +3761,32 @@ impl GpuDqnTrainer {
/// 5. If dot < 0: g_train -= (dot/|g_val|²) * g_val (project in-place)
/// 6. Swap back → grad_buf = projected g_train (zero-copy)
///
/// No DtoD copies — all data movement is through CudaSlice pointer swaps.
/// On error, caller must call `restore_grad_from_vaccine_save()` to swap back.
/// Zero-copy: g_train stays in ptrs.grad_buf, g_val is computed into
/// vaccine_grad_save scratch. Projection modifies ptrs.grad_buf in-place.
/// No CudaSlice swaps — single pointer source of truth throughout.
pub fn apply_gradient_vaccine(
&mut self,
vaccine_batch: &crate::dqn::replay_buffer_type::GpuBatch,
) -> Result<(), MLError> {
// Only run vaccine after CUDA graphs are captured (buffers initialized)
if self.graph_forward.is_none() {
return Ok(());
}
let tp = self.total_params;
let g_train_ptr = self.ptrs.grad_buf; // training grads — stays put
let g_val_ptr = self.vaccine_grad_save.raw_ptr(); // scratch for vaccine grads
// Step 1: Swap grad_buf ↔ vaccine_grad_save.
// After swap: grad_buf = empty scratch, vaccine_grad_save = g_train.
std::mem::swap(&mut self.grad_buf, &mut self.vaccine_grad_save);
let save_ptr = self.vaccine_grad_save.raw_ptr(); // g_train
let g_val_ptr = self.grad_buf.raw_ptr(); // scratch → g_val
// Step 2: Zero scratch for vaccine pass
self.stream.memset_zeros(&mut self.grad_buf)
.map_err(|e| MLError::ModelError(format!("vaccine zero grad: {e}")))?;
// Step 1: Zero scratch + d_logits for vaccine pass
self.stream.memset_zeros(&mut self.vaccine_grad_save)
.map_err(|e| MLError::ModelError(format!("vaccine zero scratch: {e}")))?;
self.stream.memset_zeros(&mut self.d_value_logits_buf)
.map_err(|e| MLError::ModelError(format!("vaccine zero d_val: {e}")))?;
self.stream.memset_zeros(&mut self.d_adv_logits_buf)
.map_err(|e| MLError::ModelError(format!("vaccine zero d_adv: {e}")))?;
// Step 3: Upload vaccine batch + forward+backward (NON-graph path)
// Step 2: Forward+backward on vaccine batch → g_val into scratch
self.upload_batch_gpu(vaccine_batch)?;
// Forward pass (cuBLAS, outside graph — reads self.grad_buf.raw_ptr() live)
self.launch_cublas_forward()?;
self.launch_curiosity_inference()?;
// Loss (C51 + MSE blend)
self.stream.memset_zeros(&mut self.d_value_logits_mse)
.map_err(|e| MLError::ModelError(format!("vaccine zero mse: {e}")))?;
self.stream.memset_zeros(&mut self.d_adv_logits_mse)
@@ -3816,28 +3799,24 @@ impl GpuDqnTrainer {
self.launch_c51_mixup()?;
self.launch_loss_reduce(&self.total_loss_buf)?;
self.launch_c51_grad()?;
// Backward (cuBLAS) — writes g_val into grad_buf (swapped scratch)
self.launch_cublas_backward()?;
// Backward writes g_val to scratch (not grad_buf)
self.launch_cublas_backward_to(g_val_ptr)?;
// Step 4: Two-phase dot product and norm² — no atomicAdd, no memset
// Step 3: dot(g_train, g_val) + ||g_val||² — two-phase reduction
let tp_i32 = tp as i32;
let blocks = self.grad_norm_blocks as u32;
let nb = blocks as i32;
let block_res_ptr = self.vaccine_block_results.raw_ptr();
let dot_norm_ptr = self.vaccine_dot_norm_buf.raw_ptr();
// Phase 1: per-block partials
unsafe {
self.stream
.launch_builder(&self.vaccine_dot_kernel)
.arg(&save_ptr)
.arg(&g_train_ptr)
.arg(&g_val_ptr)
.arg(&block_res_ptr)
.arg(&tp_i32)
.launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("vaccine dot phase1: {e}")))?;
}
// Phase 2: reduce partials
unsafe {
self.stream
.launch_builder(&self.vaccine_dot_finalize_kernel)
.arg(&block_res_ptr)
@@ -3847,12 +3826,11 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("vaccine dot phase2: {e}")))?;
}
// Step 5: Conditional projection (only when dot < 0)
// Modifies vaccine_grad_save (g_train) in-place
// Step 4: Project g_train in-place (only when dot < 0)
unsafe {
self.stream
.launch_builder(&self.vaccine_project_kernel)
.arg(&save_ptr)
.arg(&g_train_ptr)
.arg(&g_val_ptr)
.arg(&dot_norm_ptr)
.arg(&tp_i32)
@@ -3860,9 +3838,6 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("vaccine project: {e}")))?;
}
// Step 6: Swap back → grad_buf = projected g_train for Adam
std::mem::swap(&mut self.grad_buf, &mut self.vaccine_grad_save);
Ok(())
}
@@ -4028,8 +4003,8 @@ impl GpuDqnTrainer {
self.check_nan_f32(self.on_b_logits_buf.raw_ptr(), b * (b0 + b1 + b2 + b3) * na, 2)?;
// Flag 3: mse_loss_buf [1] (MSE loss scalar)
self.check_nan_f32(self.mse_loss_buf.raw_ptr(), 1, 3)?;
// Flag 6: grad_buf (cuBLAS backward output)
self.check_nan_f32(self.grad_buf.raw_ptr(), self.total_params, 6)?;
// Flag 6: grad_buf (cuBLAS backward output) — use ptrs for consistency
self.check_nan_f32(self.ptrs.grad_buf, self.total_params, 6)?;
// Flag 7: save_current_lp (softmax probs from MSE loss — f32)
self.check_nan_f32_b(self.save_current_lp.raw_ptr(), b * 4 * na, 7)?;
Ok(())
@@ -4713,10 +4688,15 @@ impl GpuDqnTrainer {
self.stream
.memset_zeros(&mut self.mse_loss_buf)
.map_err(|e| MLError::ModelError(format!("zero mse_loss: {e}")))?;
// grad_buf: backward_full uses beta=1.0 GEMM accumulation + bias atomicAdd
self.stream
.memset_zeros(&mut self.grad_buf)
.map_err(|e| MLError::ModelError(format!("zero grad_buf: {e}")))?;
// grad_buf: backward_full uses beta=1.0 GEMM accumulation — zero via ptrs
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.ptrs.grad_buf,
0,
self.total_params * std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
}
// d_value/adv_logits: c51_grad + mse_grad kernels write directly (no atomicAdd)
self.stream
.memset_zeros(&mut self.d_value_logits_buf)
@@ -4863,8 +4843,14 @@ impl GpuDqnTrainer {
.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}")))?;
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.ptrs.grad_buf,
0,
self.total_params * std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
}
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)
@@ -5771,13 +5757,17 @@ impl GpuDqnTrainer {
/// chain stays in f32 to preserve gradient precision at large batch sizes.
/// Accumulates weight gradients into `grad_buf`.
pub(crate) fn launch_cublas_backward(&self) -> Result<(), MLError> {
self.launch_cublas_backward_to(self.ptrs.grad_buf)
}
/// Backward pass writing weight gradients to an arbitrary output buffer.
/// Used by the vaccine to write g_val to scratch without touching grad_buf.
fn launch_cublas_backward_to(&self, grad_base: u64) -> Result<(), MLError> {
let bw = &self.cublas_backward;
let param_sizes = compute_param_sizes(&self.config);
let w_ptrs = f32_weight_ptrs_from_base(self.ptrs.params_ptr, &param_sizes);
let grad_base = self.grad_buf.raw_ptr(); // f32 buffer — use raw_ptr directly
// #31 Bottleneck: backward pass input is bn_concat, not raw states
let states_ptr = if self.config.bottleneck_dim > 0 {
self.bn_concat_buf.raw_ptr()

View File

@@ -1000,9 +1000,7 @@ impl FusedTrainingCtx {
if let Some(ref vaccine_batch) = self.pending_vaccine_batch {
tracing::debug!("Applying gradient vaccine at step {} for diversity correction", self.steps_since_varmap_sync);
if let Err(e) = self.trainer.apply_gradient_vaccine(vaccine_batch) {
tracing::debug!("Gradient vaccine failed, restoring grad_buf: {e}");
self.trainer.restore_grad_from_vaccine_save()
.map_err(|e2| anyhow::anyhow!("vaccine restore: {e2}"))?;
tracing::debug!("Gradient vaccine failed (non-fatal, grad_buf unchanged): {e}");
}
}
}