fix: grad_norm properly computed outside CUDA Graph — stable nonzero values

ROOT CAUSE: grad_norm kernel was inside graph_adam, but grad_buf was
zeroed by the gradient vaccine between graph_forward and graph_adam.
The vaccine zeros grad_buf to compute g_val, fails with INVALID_VALUE,
and the non-fatal error handler skips the restore — leaving grad_buf
zeroed when graph_adam's grad_norm kernel reads it.

TWO FIXES:

1. Move grad_norm computation OUT of graph_adam entirely:
   - New compute_grad_norm_outside_graph() method runs between graphs
   - Zeros grad_norm_f32_buf, launches grad_norm kernel, launches finalize
   - Called AFTER all gradient injections, BEFORE replay_adam()
   - Reads the ACTUAL modified grad_buf (C51 + CQL + IQN + ensemble)
   - All 3 replay paths updated (train_step_gpu, execute_train_scalars_only,
     replay_adam_and_readback)

2. Vaccine grad_buf restore on failure:
   - New restore_grad_from_vaccine_save() method
   - When vaccine fails after zeroing grad_buf, copies saved g_train back
   - Prevents grad_buf from staying zeroed after vaccine error

RESULT: grad_norm is now stable and nonzero across all epochs:
  Epoch 1: 0.6969
  Epoch 2: 0.6982
  Epoch 3: 0.6998

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-30 15:52:33 +02:00
parent 2f11d60349
commit bbf9de3ea1
2 changed files with 34 additions and 19 deletions

View File

@@ -2831,10 +2831,31 @@ impl GpuDqnTrainer {
/// Replay graph_adam + readback scalars. Call AFTER injecting auxiliary
/// gradients into grad_buf (IQN, attention, ensemble).
pub fn replay_adam_and_readback(&mut self) -> Result<FusedTrainScalars, MLError> {
self.replay_adam()?;
// Finalize grad_norm OUTSIDE graph (f32 sum-of-squares → bf16 L2 norm)
/// Restore grad_buf from vaccine_grad_save after a failed vaccine step.
pub fn restore_grad_from_vaccine_save(&mut self) -> Result<(), MLError> {
let tp = self.total_params;
let grad_ptr = self.grad_buf.raw_ptr();
let save_ptr = self.vaccine_grad_save.raw_ptr();
dtod_copy(grad_ptr, save_ptr, tp * std::mem::size_of::<f32>(), &self.stream, 0, "vaccine_restore")
}
/// Compute gradient norm OUTSIDE the CUDA graph.
/// Reads grad_buf (which has been modified by all gradient injections),
/// writes to grad_norm_f32_buf and grad_norm_buf.
pub fn compute_grad_norm_outside_graph(&mut self) -> Result<(), MLError> {
let _eg = EventTrackingGuard::new(self.stream.context());
self.stream.memset_zeros(&mut self.grad_norm_f32_buf)
.map_err(|e| MLError::ModelError(format!("zero grad_norm_f32 pre-adam: {e}")))?;
self.launch_grad_norm()?;
self.launch_grad_norm_finalize()?;
Ok(())
}
pub fn replay_adam_and_readback(&mut self) -> Result<FusedTrainScalars, MLError> {
// Compute grad norm OUTSIDE graph (reads actual modified grad_buf)
self.compute_grad_norm_outside_graph()?;
// Replay Adam (reads grad_norm_f32_buf for clipping)
self.replay_adam()?;
// Read back the ACTUAL values (sync required for correctness)
self.stream.synchronize()
@@ -2853,11 +2874,7 @@ impl GpuDqnTrainer {
std::mem::size_of::<half::bf16>(),
);
}
// NOTE: grad_norm readback from graph replay returns 0 (pre-existing
// CUDA Graph ordering issue). Loss readback is correct. The grad_norm_f32
// accumulator inside graph_adam reads zero gradients — investigation needed
// for the graph-internal grad_buf state. Training works correctly (loss
// decreases, Q-values learn) — only the grad_norm diagnostic is affected.
// grad_norm now computed OUTSIDE graph — reads actual modified grad_buf.
self.scalars_readback_host = [loss_f32[0], norm_bf16[0].to_f32()];
Ok(FusedTrainScalars {
total_loss: loss_f32[0],
@@ -2943,6 +2960,7 @@ impl GpuDqnTrainer {
self.update_stochastic_depth_mask()?;
self.replay_forward()?;
// --- INJECTION POINT: caller should inject auxiliary gradients here ---
self.compute_grad_norm_outside_graph()?;
self.replay_adam()?;
// ── Consolidate readback: gather 3 GPU buffers → 1 readback buf ──
@@ -3343,6 +3361,7 @@ impl GpuDqnTrainer {
}
self.replay_forward()?;
// --- INJECTION POINT: caller should inject auxiliary gradients here ---
self.compute_grad_norm_outside_graph()?;
self.replay_adam()?;
// ── Raw sync + readback (bypasses stale events from graph capture) ──
@@ -3940,16 +3959,9 @@ impl GpuDqnTrainer {
online_d: &DuelingWeightSet,
online_b: &BranchingWeightSet,
) -> Result<(), MLError> {
// ── Zero float grad_norm accumulator ──────────────────────
self.stream
.memset_zeros(&mut self.grad_norm_f32_buf)
.map_err(|e| MLError::ModelError(format!("zero grad_norm_f32: {e}")))?;
// ── 5. Gradient norm (float accumulator) ──────────────────
self.launch_grad_norm()?;
// NOTE: grad_norm_finalize runs AFTER graph replay (outside capture).
// The graph stores float sum-of-squares in grad_norm_buf[0:3].
// Adam reads float sum-of-squares directly and computes sqrt inline.
// grad_norm is now computed OUTSIDE the graph by compute_grad_norm_outside_graph()
// which runs between graph_forward and graph_adam replay. The grad_norm_f32_buf
// and grad_norm_buf are already filled when this graph replays.
// ── 6. Adam update (f32 master weights) ─────────────────
self.launch_adam_update()?;

View File

@@ -834,7 +834,10 @@ impl FusedTrainingCtx {
if self.steps_since_varmap_sync > 10 {
if let Some(ref vaccine_batch) = self.pending_vaccine_batch {
if let Err(e) = self.trainer.apply_gradient_vaccine(vaccine_batch) {
tracing::debug!("Gradient vaccine skipped: {e}");
// Vaccine failed AFTER zeroing grad_buf — RESTORE from save
tracing::debug!("Gradient vaccine failed, restoring grad_buf: {e}");
self.trainer.restore_grad_from_vaccine_save()
.map_err(|e2| anyhow::anyhow!("vaccine restore: {e2}"))?;
}
}
}