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:
@@ -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, ¶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()
|
||||
|
||||
@@ -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}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user