From 4159d201a027ba765bb3810d2313db79beef08ec Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 13 Apr 2026 23:05:43 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20eliminate=20grad=5Fbuf=20pointer=20dupli?= =?UTF-8?q?cation=20=E2=80=94=20avg=5Fgrad=3D0=20root=20cause?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 94 +++++++++---------- crates/ml/src/trainers/dqn/fused_training.rs | 4 +- 2 files changed, 43 insertions(+), 55 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 327dba77d..a89c2fb7d 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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::(), + 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::(), + 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, ¶m_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() diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index d77620680..58f658811 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -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}"); } } }