diag(ml-alpha): scan Mamba2 L1 fwd between projections and SSM kernel

Add labels 43-46 inside Mamba2Block::forward_train_seq_into, scanned
BETWEEN the three projection cuBLAS GEMMs (w_in, w_a, w_b) and the
SSM scan kernel launch (mamba2_alpha_scan_fwd_seq). The scan kernel
ONLY reads a_proj/b_proj/w_c/h_s2 and ONLY writes h_enriched_seq -
it cannot retroactively corrupt its read-only inputs, so a NaN
observed here pinpoints the projection cuBLAS path (PROJ verdict)
vs the scan kernel itself (SCAN verdict, requires labels 43-46
clean AND existing labels 36/37 still firing).

Plumbed via the existing NanScanHook installed from IntegratedTrainer::new
on mamba2_l1 only (L2 left unhooked - scope is the L1 forward path
established as the failure window by labels 36/37). Hook source
identical to the perception trainer's hook (shared cubin handle +
ISV step counter pointer); per-launch dispatch is a no-op when
FOXHUNT_NAN_SCAN is unset, so production training pays zero cost.

  43 - mamba2_l1_fwd_a_proj_pre_scan (post w_a, pre scan)
  44 - mamba2_l1_fwd_b_proj_pre_scan (post w_b, pre scan)
  45 - mamba2_l1_fwd_x_pre_scan      (post w_in)
  46 - mamba2_l1_fwd_h_s2_pre_scan   (zero-init residual sanity)

Refs pearl_atomicadd_masks_v_instability.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-29 10:26:48 +02:00
parent 85531c8c3a
commit fb346a3cad
3 changed files with 130 additions and 0 deletions

View File

@@ -129,6 +129,36 @@
// comparison probe to 41 — if 42 NaN but
// 41 clean, corruption is L2-side first)
//
// Mamba2 L1 forward sub-stage scans (PROJ vs SCAN disambiguation —
// scanned INSIDE Mamba2Block::forward_train_seq_into between the three
// projection cuBLAS GEMMs and the SSM scan kernel launch). The scan
// kernel (`mamba2_alpha_scan_fwd_seq`) reads a_proj/b_proj/w_c/h_s2 and
// ONLY writes h_enriched_seq — it cannot retroactively corrupt its own
// inputs, so a NaN observed here pinpoints the projection cuBLAS path.
// Pairs with labels 36/37 (same buffers, scanned in PerceptionTrainer
// IMMEDIATELY after forward_train_seq_into returns).
//
// 43 — mamba2_l1_fwd_a_proj_pre_scan (post-w_a, pre-scan; matches
// label 36/20 buffer content)
// 44 — mamba2_l1_fwd_b_proj_pre_scan (post-w_b, pre-scan; the
// second scan input, never had
// its own label before)
// 45 — mamba2_l1_fwd_x_pre_scan (post-w_in; the input to both
// w_a and w_b gemms; matches
// label 37/21 buffer content)
// 46 — mamba2_l1_fwd_h_s2_pre_scan (zero-init recurrent residual;
// sanity probe for upstream
// corruption of a buffer that
// should always be exactly zero)
//
// Verdict mapping:
// any of {43, 44, 45} NaN at step 4 → PROJ verdict (projection
// cuBLAS gemms or vsn_out_d input produced NaN); next bisection
// step is vsn_out_d (label 6) and the input weights (labels 41+).
// {43, 44, 45, 46} CLEAN at step 4 but {36, 37} fire → SCAN verdict
// (impossible under static reading of the scan kernel — escalate
// to a Heisenbug/UPSTREAM investigation).
//
// Format printed on hit (one line per offending block):
// "NAN_SCAN step=<step> label=<name> first_bad_idx=<i> val=<f>\n"
//
@@ -235,6 +265,10 @@ extern "C" __global__ void nan_scan(
case 40: label = "step_entry_vsn_b"; break;
case 41: label = "step_entry_mamba2_l1_w_a"; break;
case 42: label = "step_entry_mamba2_l2_w_a"; break;
case 43: label = "mamba2_l1_fwd_a_proj_pre_scan"; break;
case 44: label = "mamba2_l1_fwd_b_proj_pre_scan"; break;
case 45: label = "mamba2_l1_fwd_x_pre_scan"; break;
case 46: label = "mamba2_l1_fwd_h_s2_pre_scan"; break;
default: break;
}
const int step = (int)isv[RL_STEP_COUNTER_ISV_INDEX];

View File

@@ -485,6 +485,16 @@ pub struct Mamba2Block {
/// per-call alloc; the buffer stays alive for the lifetime of the
/// block, large enough for all gemm sizes used by the projections.
_cublas_workspace: CudaSlice<u8>,
/// Diagnostic NaN-scan hook for the step-4 NaN hunt (2026-05-29
/// follow-up to `pearl_atomicadd_masks_v_instability`). Mirrors the
/// hook plumbed into `PerceptionTrainer` so the Mamba2 forward path
/// can scan sub-stage outputs (post-projection-gemms / pre-scan-
/// kernel) BETWEEN the projection cuBLAS calls and the SSM scan
/// kernel launch. `None` outside the alpha-rl training entrypoint —
/// supervised path leaves the hook uninstalled so its forward pays
/// zero per-launch cost. Set via [`install_nan_scan_hook`].
nan_scan_hook: Option<crate::trainer::perception::NanScanHook>,
}
impl Mamba2Block {
@@ -601,9 +611,52 @@ impl Mamba2Block {
kernel_adamw,
cublas,
_cublas_workspace: cublas_workspace,
nan_scan_hook: None,
})
}
/// Install the shared diagnostic NaN-scan hook. Wired from
/// `IntegratedTrainer::new` for the alpha-rl training path so
/// `forward_train_seq_into` can scan post-projection / pre-scan
/// sub-stages (labels 4346 in the canonical table — see
/// `cuda/nan_scan.cu`). Hook holds the kernel function handle, the
/// ISV device pointer for the step counter, and the FOXHUNT_NAN_SCAN
/// gate; per-launch dispatch is a no-op when the env var is unset.
pub fn install_nan_scan_hook(&mut self, hook: crate::trainer::perception::NanScanHook) {
self.nan_scan_hook = Some(hook);
}
/// Per-sub-stage NaN scan inside the Mamba2 forward path. Mirrors
/// `PerceptionTrainer::nan_scan` byte-for-byte (single block × 32
/// threads, no shared mem, no atomics; ISV step counter read inside
/// the kernel). Hook source instead of `self` for the kernel + ISV
/// pointer. No-op when no hook is installed or the env var is unset.
fn nan_scan(&self, label_id: i32, buf_ptr: u64, n: usize) -> Result<()> {
let hook = match self.nan_scan_hook.as_ref() {
Some(h) if h.enabled => h,
_ => return Ok(()),
};
let n_i = n as i32;
let mut args = RawArgs::new();
args.push_ptr(buf_ptr);
args.push_i32(n_i);
args.push_i32(label_id);
args.push_ptr(hook.isv_dev_ptr);
let mut ptrs = args.build_arg_ptrs();
unsafe {
raw_launch(
hook.fn_handle.cu_function(),
(1, 1, 1),
(32, 1, 1),
0,
self.raw_stream,
&mut ptrs[..args.len()],
)
.map_err(|e| anyhow!("mamba2 nan_scan(label={}): {:?}", label_id, e))?;
}
Ok(())
}
// ── Forward pass ───────────────────────────────────────────────────
/// GPU-native forward pass with full activation cache returned alongside
@@ -1248,6 +1301,35 @@ impl Mamba2Block {
&self.cublas, &self.stream, &mut scratch.b_proj,
).map_err(|e| anyhow!("w_b fwd_into: {e}"))?;
// Step-4 NaN localisation (2026-05-29 follow-up to
// pearl_atomicadd_masks_v_instability). Scan the post-projection
// sub-stage buffers BEFORE the scan kernel runs. The scan kernel
// (`mamba2_alpha_scan_fwd_seq`) ONLY reads a_proj / b_proj /
// w_c / h_s2 and ONLY writes h_enriched_seq — it cannot
// retroactively corrupt its read-only inputs. So a NaN observed
// here pinpoints the projection cuBLAS path (or the vsn_out_d
// input flowing into the w_in gemm) as the source.
//
// 43 — a_proj after w_a gemm; INPUT TO SCAN
// 44 — b_proj after w_b gemm; INPUT TO SCAN
// 45 — x after w_in gemm; INPUT TO w_a / w_b gemms
// 46 — h_s2 initial recurrent state (zero-init from
// construction; sanity check that it has not
// been corrupted by an upstream write)
//
// Verdict mapping (paired with existing labels 36/37 post-call):
// any of {43, 44, 45} fires → projection produced NaN; bisect
// further by checking vsn_out_d (label 6 at PerceptionTrainer
// forward_only_dispatch) or the input gemm weights themselves
// (label 41 at top-of-step persistent weights).
// {43, 44, 45, 46} all clean but {36, 37} fire at perception
// post-call → impossible under static reading of the scan
// kernel; escalate to UPSTREAM.
self.nan_scan(43, scratch.a_proj.cuda_data().raw_ptr(), scratch.a_proj.numel())?;
self.nan_scan(44, scratch.b_proj.cuda_data().raw_ptr(), scratch.b_proj.numel())?;
self.nan_scan(45, scratch.x.cuda_data().raw_ptr(), scratch.x.numel())?;
self.nan_scan(46, scratch.h_s2.cuda_data().raw_ptr(), scratch.h_s2.numel())?;
// 4. scan_fwd_seq → scratch.h_enriched_seq.
// h_s2 stays zero from construction (never written in supervised path).
let block_threads: u32 = 32;

View File

@@ -2870,6 +2870,20 @@ impl IntegratedTrainer {
},
);
// Mamba2 L1 forward sub-stage NaN scans (labels 4346, added in
// step-4 NaN localisation follow-up to pearl_atomicadd_masks_v_
// instability). Hook installed only on L1 because the established
// failure window is L1's forward path (labels 36/37 fire NaN at
// save time — see pearl entry). L2's forward gets no hook; the
// diagnostic is scoped to L1 to keep noise minimal.
trainer.perception.trunk.mamba2_l1_mut().install_nan_scan_hook(
crate::trainer::perception::NanScanHook {
fn_handle: trainer.nan_scan_fn.clone(),
isv_dev_ptr: trainer.isv_dev_ptr,
enabled: trainer.nan_scan_enabled,
},
);
Ok(trainer.with_controllers_bootstrapped()?)
}