fix(nan-diag): wire NaN checks into captured graph + fix stale labels

Fold 1 of train-multi-seed-72fl6 hit NaN at step 5 with `flagged=[]` —
diagnostic told us nothing because:

1. The 8 NaN-check kernels in `run_nan_checks_*_forward` were never
   invoked from production training. Allocated-zero flags read back
   as zeros; halt-on-NaN reported "no buffer flagged" while loss was
   demonstrably NaN via host-side pinned readback.
2. The error message labels (training_loop.rs:1853) referenced
   `bf16_params` — dead code from before the bf16→TF32 switch — and
   the format string did not match the actual kernel index→buffer
   mapping.

Wired the existing NaN check kernels into `submit_post_aux_ops` so
they run every step inside the captured parent graph (post_aux child).
Flags persist across steps within a fold so the FIRST buffer to go
NaN remains visible at host-readback time; reset is host-side at fold
boundary in `reset_for_fold`.

Extended coverage beyond the original 8 buffers (states, on_v_logits,
on_b_logits, mse_loss, params, grad_buf, save_current_lp) with 4
loss-component buffers most likely to break under post-S&P + adversarial
regime stress:

  Flag 8  save_projected      C51 target distribution
  Flag 9  moe_gate_softmax    MoE gate over 8 experts
  Flag 10 aux_nb_loss_scalar  aux next-bar MSE
  Flag 11 aux_rg_loss_scalar  aux regime CE

Flag buffer grown 8 → 16; slots 12-15 reserved for future expansion
(CQL / IQN / per-branch backward).

Updated training_loop.rs error message: 16 named slots matching actual
kernel layout, per-flag indexed name in flagged list, dropped stale
`bf16_params` / `_bf16` suffixes, explicit hint when `flagged=[]`.

Cost per step: ~12 single-block reductions × ~5µs = ~60µs. <0.1%
overhead at the 258ms/step measured on L40S — pinpoints source
buffer on next NaN halt.

Per `feedback_no_legacy_aliases.md`, `feedback_no_stubs.md`,
`feedback_trust_code_not_docs.md`.
This commit is contained in:
jgrusewski
2026-04-28 11:24:22 +02:00
parent 391f6d8d58
commit 756b1ef317
4 changed files with 126 additions and 13 deletions

View File

@@ -10656,7 +10656,7 @@ impl GpuDqnTrainer {
aux_knb, aux_kr, aux_h, max_aux_tensor_len,
);
let nan_flags_buf = stream.alloc_zeros::<i32>(8)
let nan_flags_buf = stream.alloc_zeros::<i32>(16)
.map_err(|e| MLError::ModelError(format!("nan_flags alloc: {e}")))?;
// v8: PopArt running statistics buffers
@@ -14247,9 +14247,12 @@ impl GpuDqnTrainer {
/// Run pre-forward NaN checks: f32 params_buf (flag 4) + params_ptr alias (flag 5).
/// Detects if previous step's Adam corrupted the weights.
///
/// Does NOT reset flags — flags persist across steps so the FIRST buffer to go
/// NaN remains visible at host-readback time. Caller resets flags only at fold
/// boundary via `reset_nan_flags()`.
pub fn run_nan_checks_pre_forward(&mut self) -> Result<(), MLError> {
let tp = self.total_params;
self.reset_nan_flags()?;
// Flag 4: f32 params_buf
self.check_nan_f32(self.params_buf.raw_ptr(), tp, 4)?;
// Flag 5: params_ptr (same buffer, second check for redundancy)
@@ -14258,7 +14261,23 @@ impl GpuDqnTrainer {
}
/// Run all post-forward NaN checks (GPU-side, no CPU sync).
/// Checks cuBLAS output logits, MSE intermediates, gradients, and params.
/// Checks cuBLAS output logits, MSE intermediates, gradients, params, and
/// loss-component buffers (C51 target dist, MoE gate softmax, aux losses).
///
/// Flag index map (16 slots):
/// 0 states_buf (f32 padded) — input
/// 1 on_v_logits — cuBLAS forward (value stream)
/// 2 on_b_logits — cuBLAS forward (branch advantage)
/// 3 mse_loss_buf [1] — MSE loss scalar
/// 4 params_buf (f32, pre-forward) — set by run_nan_checks_pre_forward
/// 5 params_ptr (redundant pre-fwd) — set by run_nan_checks_pre_forward
/// 6 grad_buf — cuBLAS backward output
/// 7 save_current_lp — C51 online softmax probs (post-fwd)
/// 8 save_projected — C51 target distribution (most NaN-prone)
/// 9 moe_gate_softmax — MoE gate softmax (saturates → NaN)
/// 10 aux_nb_loss_scalar [1] — aux next-bar MSE loss
/// 11 aux_rg_loss_scalar [1] — aux regime CE loss
/// 12-15 reserved (CQL / IQN / future)
pub fn run_nan_checks_post_forward(&mut self, batch_size: usize) -> Result<(), MLError> {
let b = batch_size;
let na = self.config.num_atoms;
@@ -14266,7 +14285,7 @@ impl GpuDqnTrainer {
let b1 = self.config.branch_1_size;
let b2 = self.config.branch_2_size;
let b3 = self.config.branch_3_size;
// Don't reset — pre-forward already set flags 4-5, we add 0-3
// Don't reset — pre-forward already set flags 4-5, we add 0-3 + 6-11.
// Flag 0: states_buf (f32 padded — if states have NaN, everything downstream does)
let pad_sd = ml_core::state_layout::STATE_DIM_PADDED;
self.check_nan_f32(self.states_buf.raw_ptr(), b * pad_sd, 0)?;
@@ -14278,8 +14297,23 @@ impl GpuDqnTrainer {
self.check_nan_f32(self.mse_loss_dev_ptr, 1, 3)?;
// 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)
// Flag 7: save_current_lp (C51 online softmax probs — f32)
self.check_nan_f32_b(self.save_current_lp.raw_ptr(), b * 4 * na, 7)?;
// Flag 8: save_projected (C51 target distribution — most NaN-prone under
// adversarial reward stress; categorical projection of TD targets onto atoms)
self.check_nan_f32(self.save_projected.raw_ptr(), b * 4 * na, 8)?;
// Flag 9: moe_gate_softmax (gate distribution over 8 experts — softmax
// saturation under post-S&P weights produces NaN that propagates into the
// MoE mixture forward pass)
self.check_nan_f32(
self.moe_gate_softmax_buf.raw_ptr(),
b * MOE_NUM_EXPERTS,
9,
)?;
// Flag 10: aux_nb_loss_scalar [1] (aux next-bar MSE loss)
self.check_nan_f32(self.aux_nb_loss_scalar_buf.raw_ptr(), 1, 10)?;
// Flag 11: aux_rg_loss_scalar [1] (aux regime CE loss)
self.check_nan_f32(self.aux_rg_loss_scalar_buf.raw_ptr(), 1, 11)?;
Ok(())
}
@@ -14327,9 +14361,10 @@ impl GpuDqnTrainer {
Ok(())
}
/// Read NaN flags back to CPU (synchronizes stream). Returns [8] flags.
pub fn read_nan_flags(&self) -> Result<[i32; 8], MLError> {
let mut host = [0_i32; 8];
/// Read NaN flags back to CPU (synchronizes stream). Returns [16] flags.
/// Index → buffer mapping is documented in `run_nan_checks_post_forward`.
pub fn read_nan_flags(&self) -> Result<[i32; 16], MLError> {
let mut host = [0_i32; 16];
self.stream.memcpy_dtoh(&self.nan_flags_buf, &mut host)
.map_err(|e| MLError::ModelError(format!("nan_flags read: {e}")))?;
Ok(host)
@@ -16205,6 +16240,15 @@ impl GpuDqnTrainer {
// Multi-horizon value forward (degenerate DtoD copy)
self.multi_horizon_value_forward(batch_size)?;
// NaN diagnostic — runs every step, flags persist across steps so the
// first buffer to go NaN remains visible when host-side guard halts.
// Reset is host-side at fold boundary (`reset_nan_flags()` from
// `reset_for_fold`). Pre-forward checks (flags 4-5) check params_buf
// which is updated by adam; checking here = checking the params Adam
// just produced (== params the NEXT step's forward will use).
self.run_nan_checks_pre_forward()?;
self.run_nan_checks_post_forward(batch_size)?;
Ok(())
}

View File

@@ -921,6 +921,13 @@ impl FusedTrainingCtx {
self.cached_median = 0.0;
self.cached_iqr = 0.0;
}
// Reset NaN diagnostic flags so previous fold's NaN events don't taint
// the next fold's diagnostic. Per-step checks inside the captured graph
// (`submit_post_aux_ops`) accumulate flags across steps within a fold so
// the FIRST buffer to go NaN remains visible at host-readback time.
self.trainer.reset_nan_flags()
.map_err(|e| anyhow::anyhow!("nan_flags reset at fold boundary: {e}"))?;
Ok(())
}
@@ -3318,7 +3325,9 @@ impl FusedTrainingCtx {
}
/// Read NaN detection flags (synchronizes stream). Used by training guard on NaN halt.
pub(crate) fn read_nan_flags(&self) -> Result<[i32; 8], crate::MLError> {
/// Index → buffer mapping is documented in
/// `GpuDqnTrainer::run_nan_checks_post_forward`.
pub(crate) fn read_nan_flags(&self) -> Result<[i32; 16], crate::MLError> {
self.trainer.read_nan_flags()
}

View File

@@ -1844,13 +1844,31 @@ impl DQNTrainer {
if gr.halt_nan {
if let Some(ref mut fused) = self.fused_ctx {
if let Ok(flags) = fused.read_nan_flags() {
let names = ["STATES_bf16", "on_v_logits", "on_b_logits", "mse_loss", "f32_params_PRE", "bf16_params_PRE", "grad_buf", "save_probs_bf16"];
let flagged: Vec<_> = flags.iter().enumerate()
// Index → buffer name. Must match
// GpuDqnTrainer::run_nan_checks_post_forward and
// run_nan_checks_pre_forward. 16-slot layout.
let names = [
"states_buf", // 0 input states (post-upload)
"on_v_logits", // 1 cuBLAS fwd: value stream
"on_b_logits", // 2 cuBLAS fwd: branch advantage
"mse_loss_scalar", // 3 MSE loss [1]
"params_buf_pre_fwd", // 4 f32 params (pre-forward)
"params_ptr_pre_fwd", // 5 same buffer (redundant pre-fwd)
"grad_buf", // 6 cuBLAS bwd output
"save_current_lp", // 7 C51 online softmax probs
"save_projected", // 8 C51 target distribution
"moe_gate_softmax", // 9 MoE gate over 8 experts
"aux_nb_loss_scalar", // 10 aux next-bar MSE [1]
"aux_rg_loss_scalar", // 11 aux regime CE [1]
"rsv12", "rsv13", "rsv14", "rsv15", // reserved
];
let flagged: Vec<String> = flags.iter().enumerate()
.filter(|(_, &f)| f != 0)
.map(|(i, _)| names[i])
.map(|(i, _)| format!("{}={}", i, names[i]))
.collect();
tracing::error!(
"NaN SOURCE at step {}: flagged=[{}] (0=mse_loss 1=grad_buf 2=d_val_logits 3=bf16_params)",
"NaN SOURCE at step {}: flagged=[{}] (empty = NaN entered via \
loss-component buffer not in checks 0..11; expand coverage)",
train_step_count, flagged.join(", ")
);
}

View File

@@ -2,6 +2,48 @@
**Status:** Populated during Plan 1 Task 6 (A.5 orphan audit). Updated on every commit per Invariant 7.
NaN diagnostic wire-up + label fix (2026-04-28): fold 1 of
`train-multi-seed-72fl6` hit NaN at step 5 with `flagged=[]` because the
8 NaN-check kernels in `run_nan_checks_pre_forward` /
`run_nan_checks_post_forward` (gpu_dqn_trainer.rs:14250+, 14278+) were
allocated but never invoked from production training — orphan
diagnostic infrastructure. The `nan_flags_buf` always read back as
zeros while host-side pinned `total_loss` confirmed NaN. Compounding
this, the error message at `training_loop.rs:1853` referenced
`bf16_params` (dead code from before the bf16→TF32 switch) and the
format-string indices did not match the actual kernel index→buffer
mapping. **Fix:** wire both pre/post NaN checks into
`submit_post_aux_ops` (captured in the `post_aux` child graph, runs
every step inside the parent `cuGraphLaunch`); flags persist across
steps within a fold so the FIRST buffer to go NaN remains visible at
host-readback time; reset is host-side at fold boundary in
`reset_for_fold`. Coverage extended from 8 → 12 active slots (16
allocated total): added flag 8 = `save_projected` (C51 target
distribution — categorical projection of TD targets, NaN-prone under
adversarial reward stress), flag 9 = `moe_gate_softmax` (gate softmax
saturation under post-S&P weights), flag 10 = `aux_nb_loss_scalar`,
flag 11 = `aux_rg_loss_scalar`. Slots 12-15 reserved (CQL / IQN /
per-branch backward). `read_nan_flags()` return type grown
`[i32; 8]``[i32; 16]` in both `GpuDqnTrainer` and
`FusedTrainingCtx`. `run_nan_checks_pre_forward` no longer resets
flags (caller's responsibility — fold-boundary host call). Updated
`training_loop.rs` error message: 16 named slots matching actual
kernel layout, per-flag indexed name in flagged list (e.g.
`[8=save_projected]`), dropped all stale `bf16_params` / `_bf16`
suffixes, explicit hint when `flagged=[]` (NaN entered via
still-uncovered buffer; expand coverage). Cost per step: ~12
single-block NaN-check reductions × ~5µs = ~60µs (<0.1% of the
258ms/step measured on L40S). Honours `feedback_no_legacy_aliases.md`
(drop bf16_params), `feedback_no_stubs.md` (orphan infrastructure
either wired or deleted — wiring chosen), and
`feedback_trust_code_not_docs.md` (label/comment was stale; verified
against actual kernel ordering). Touched: `gpu_dqn_trainer.rs`
(`nan_flags_buf` 8→16, `read_nan_flags` return type, no-reset
pre-forward, +4 post-forward checks, `submit_post_aux_ops` adds both
NaN checks); `fused_training.rs` (return type, fold-boundary
`reset_nan_flags()` call); `training_loop.rs` (error message
overhaul). cargo check clean at 13 warnings (workspace baseline).
TLOB cuBLAS-Lt → classic `cublasSgemm_v2` migration (2026-04-28):
followup to the symbol-rename fix below. The cuBLAS-Lt heuristic
`cublasLtMatmulAlgoGetHeuristic` returns "no algo found" for TLOB's