diag: GPU-side NaN detection kernels + capture warmup investigation

Added dqn_nan_check_f32/bf16 kernels for zero-overhead GPU-side NaN
detection between graph stages. Key findings:

- NaN appears in f32 params at step 0 (before first training forward)
- States (input data) are always clean
- The graph capture warmup launches corrupt params via:
  C51 forward → NaN gradients → Adam warmup → NaN weights
- Removed Adam warmup launch, but NaN persists — the forward
  graph warmup itself may corrupt buffers through bf16 overflow
  in cuBLAS backward path
- Step 14/22 NaN in later runs is from accumulated corruption

Next: investigate whether cuBLAS backward with random Xavier weights
produces extreme gradients that overflow bf16 d_logits staging buffers.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-05 00:39:34 +02:00
parent b316a25262
commit 57844e1cc5
4 changed files with 195 additions and 9 deletions

View File

@@ -262,6 +262,46 @@ extern "C" __global__ void dqn_clip_grad_kernel(
* Launch config: grid=(ceil(n/256), 1, 1), block=(256, 1, 1).
* ══════════════════════════════════════════════════════════════════════ */
/* ══════════════════════════════════════════════════════════════════════
* NaN DETECTION KERNEL
*
* Scans a float buffer for NaN/Inf values. Writes 1 to flags[flag_idx]
* if ANY non-finite value is found. Used for diagnostic instrumentation
* between CUDA graph replays to pinpoint NaN source.
*
* Launch config: grid=(ceil(n/256), 1, 1), block=(256, 1, 1).
* ══════════════════════════════════════════════════════════════════════ */
extern "C" __global__ void dqn_nan_check_f32(
const float* __restrict__ buf,
int n,
int* __restrict__ flags,
int flag_idx
) {
for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += gridDim.x * blockDim.x) {
if (!isfinite(buf[i])) {
flags[flag_idx] = 1;
return;
}
}
}
/* BF16 variant for checking bf16 buffers (cuBLAS outputs, weight shadows) */
extern "C" __global__ void dqn_nan_check_bf16(
const __nv_bfloat16* __restrict__ buf,
int n,
int* __restrict__ flags,
int flag_idx
) {
for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += gridDim.x * blockDim.x) {
float v = (float)buf[i];
if (!isfinite(v)) {
flags[flag_idx] = 1;
return;
}
}
}
extern "C" __global__ void dqn_zero_kernel(
__nv_bfloat16* __restrict__ buf,
int n

View File

@@ -662,6 +662,11 @@ pub struct GpuDqnTrainer {
pruning_compute_kernel: CudaFunction,
/// HER in-place goal relabel kernel (writes directly into padded staging buffers).
pub(crate) her_inplace_kernel: CudaFunction,
/// NaN detection kernels (GPU-side, no CPU readback per step).
pub(crate) nan_check_f32_kernel: CudaFunction,
pub(crate) nan_check_bf16_kernel: CudaFunction,
/// NaN flags buffer [8] — one flag per checkpoint. Reset per step, read at epoch boundary.
pub(crate) nan_flags_buf: CudaSlice<i32>,
/// #20 Pruning epoch (epoch at which to compute the mask).
pruning_epoch: usize,
/// #20 Pruning fraction (0.7 = prune 70% of smallest weights).
@@ -2201,7 +2206,7 @@ impl GpuDqnTrainer {
// per array. Stack is set once in DQNTrainer::new() (64KB for all kernels).
// ── Compile 4 utility kernels (grad_norm, adam_update, BF16 converters) ─
let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, scale_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, causal_intervene_kernel_fn, causal_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel) =
let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, scale_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, causal_intervene_kernel_fn, causal_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel, nan_check_f32_kernel, nan_check_bf16_kernel) =
compile_training_kernels(&stream, &config)?;
// Separate grad_norm instance for non-graph launches (clip_grad_buf_inplace).
@@ -2690,6 +2695,8 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("alloc bn_d_concat: {e}")))?;
let bn_d_hidden_buf = stream.alloc_zeros::<f32>(b * bn_alloc_dim)
.map_err(|e| MLError::ModelError(format!("alloc bn_d_hidden: {e}")))?;
let nan_flags_buf = stream.alloc_zeros::<i32>(8)
.map_err(|e| MLError::ModelError(format!("nan_flags alloc: {e}")))?;
Ok(Self {
config,
stream,
@@ -2844,6 +2851,9 @@ impl GpuDqnTrainer {
pruning_mask_kernel,
pruning_compute_kernel,
her_inplace_kernel,
nan_check_f32_kernel,
nan_check_bf16_kernel,
nan_flags_buf,
pruning_epoch: prune_ep,
pruning_fraction: prune_frac,
causal_intervene_kernel: causal_intervene_kernel_fn,
@@ -3645,6 +3655,88 @@ impl GpuDqnTrainer {
&self.td_errors_buf
}
/// Run pre-forward NaN checks: bf16 params (flag 4) + f32 master params (flag 5).
/// Detects if previous step's Adam corrupted the weights.
pub fn run_nan_checks_pre_forward(&mut self) -> Result<(), MLError> {
let tp = self.total_params;
self.reset_nan_flags()?;
// Flag 4: f32 master params
self.check_nan_f32(self.params_buf.raw_ptr(), tp, 4)?;
// Flag 5: bf16 shadow params (used by cuBLAS forward)
self.check_nan_bf16(self.ptrs.params_buf, tp, 5)?;
Ok(())
}
/// Run all post-forward NaN checks (GPU-side, no CPU sync).
/// Checks cuBLAS output logits, MSE intermediates, gradients, and params.
pub fn run_nan_checks_post_forward(&mut self, batch_size: usize) -> Result<(), MLError> {
let b = batch_size;
let na = self.config.num_atoms;
let b0 = self.config.branch_0_size;
let b1 = self.config.branch_1_size;
let b2 = self.config.branch_2_size;
// Don't reset — pre-forward already set flags 4-5, we add 0-3
// Flag 0: states_buf (bf16 padded — if states have NaN, everything downstream does)
let pad_sd = (self.config.state_dim + 127) & !127;
self.check_nan_bf16(self.states_buf.raw_ptr(), b * pad_sd, 0)?;
// Flag 1: on_v_logits (f32 cuBLAS forward output — value stream logits)
self.check_nan_f32(self.on_v_logits_buf.raw_ptr(), b * na, 1)?;
// Flag 2: on_b_logits (f32 cuBLAS forward output — branch advantage logits)
self.check_nan_f32(self.on_b_logits_buf.raw_ptr(), b * (b0 + b1 + b2) * 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 7: save_current_lp (softmax probs from MSE loss — bf16)
self.check_nan_bf16(self.save_current_lp.raw_ptr(), b * 3 * na, 7)?;
Ok(())
}
/// Launch GPU-side NaN check on an f32 buffer. Writes 1 to nan_flags_buf[flag_idx] if NaN found.
/// No CPU sync — stays entirely on GPU. Call reset_nan_flags() before a batch of checks.
pub fn check_nan_f32(&self, buf_ptr: u64, n: usize, flag_idx: i32) -> Result<(), MLError> {
let n_i32 = n as i32;
let blocks = ((n + 255) / 256) as u32;
let flags_ptr = self.nan_flags_buf.raw_ptr();
unsafe {
self.stream
.launch_builder(&self.nan_check_f32_kernel)
.arg(&buf_ptr).arg(&n_i32).arg(&flags_ptr).arg(&flag_idx)
.launch(LaunchConfig { grid_dim: (blocks.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("nan_check_f32[{flag_idx}]: {e}")))?;
}
Ok(())
}
/// Launch GPU-side NaN check on a bf16 buffer.
pub fn check_nan_bf16(&self, buf_ptr: u64, n: usize, flag_idx: i32) -> Result<(), MLError> {
let n_i32 = n as i32;
let blocks = ((n + 255) / 256) as u32;
let flags_ptr = self.nan_flags_buf.raw_ptr();
unsafe {
self.stream
.launch_builder(&self.nan_check_bf16_kernel)
.arg(&buf_ptr).arg(&n_i32).arg(&flags_ptr).arg(&flag_idx)
.launch(LaunchConfig { grid_dim: (blocks.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("nan_check_bf16[{flag_idx}]: {e}")))?;
}
Ok(())
}
/// Zero the NaN flags buffer (call before a batch of check_nan calls).
pub fn reset_nan_flags(&mut self) -> Result<(), MLError> {
self.stream.memset_zeros(&mut self.nan_flags_buf)
.map_err(|e| MLError::ModelError(format!("nan_flags reset: {e}")))
}
/// 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];
self.stream.memcpy_dtoh(&self.nan_flags_buf, &mut host)
.map_err(|e| MLError::ModelError(format!("nan_flags read: {e}")))?;
Ok(host)
}
/// Cast f32 source buffer → bf16 td_errors_buf (for IQN f32 → PER bf16 boundary).
pub fn cast_f32_to_td_errors(&self, src: &CudaSlice<f32>) -> Result<(), MLError> {
let n = self.config.batch_size as i32;
@@ -4110,15 +4202,19 @@ impl GpuDqnTrainer {
// Launch both graphs on first capture (warm-up). On subsequent steps,
// only graph_forward is replayed by train_step_gpu(). The caller then
// injects auxiliary gradients and calls replay_adam_and_readback().
graph_fwd.launch().map_err(|e| {
MLError::ModelError(format!("CUDA graph_forward first launch: {e}"))
})?;
// Launch forward graphs for warmup ONLY — no Adam.
// With random Xavier weights, the first forward+backward produces extreme
// gradients. Running Adam on these corrupts the initial weights with NaN.
// The first real training step runs Adam with proper gradient clipping.
graph_mse.launch().map_err(|e| {
MLError::ModelError(format!("CUDA graph_forward_mse first launch: {e}"))
})?;
graph_adam.launch().map_err(|e| {
MLError::ModelError(format!("CUDA graph_adam first launch: {e}"))
graph_fwd.launch().map_err(|e| {
MLError::ModelError(format!("CUDA graph_forward first launch: {e}"))
})?;
// Zero grad_buf after C51 warmup to prevent NaN from leaking into step 0
self.stream.memset_zeros(&mut self.grad_buf)
.map_err(|e| MLError::ModelError(format!("zero grad_buf post-capture: {e}")))?;
info!(
"GpuDqnTrainer: 3 CUDA graphs captured and launched \
@@ -5755,7 +5851,7 @@ impl GpuDqnTrainer {
fn compile_training_kernels(
stream: &Arc<CudaStream>,
config: &GpuDqnTrainConfig,
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
info!(
state_dim = config.state_dim,
total_params = compute_total_params(config),
@@ -5824,9 +5920,13 @@ fn compile_training_kernels(
let her_inplace = module.load_function("her_inplace_relabel")
.map_err(|e| MLError::ModelError(format!("her_inplace_relabel load: {e}")))?;
let nan_check_f32 = module.load_function("dqn_nan_check_f32")
.map_err(|e| MLError::ModelError(format!("dqn_nan_check_f32 load: {e}")))?;
let nan_check_bf16 = module.load_function("dqn_nan_check_bf16")
.map_err(|e| MLError::ModelError(format!("dqn_nan_check_bf16 load: {e}")))?;
info!("GpuDqnTrainer: 28 utility kernels loaded from precompiled cubin");
Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32, scale_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, causal_intervene, causal_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace))
info!("GpuDqnTrainer: 30 utility kernels loaded from precompiled cubin");
Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32, scale_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, causal_intervene, causal_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace, nan_check_f32, nan_check_bf16))
}
/// Load the standalone Polyak EMA kernel from precompiled cubin.

View File

@@ -655,10 +655,34 @@ impl FusedTrainingCtx {
let gpu_batch = batch.gpu_batch.as_ref()
.ok_or_else(|| anyhow::anyhow!("Fused training requires gpu_batch (GPU PER)"))?;
// ── NaN detection: early steps with per-stage checks ────────────
let nan_diag = self.steps_since_varmap_sync < 20;
if nan_diag {
self.trainer.run_nan_checks_pre_forward()?;
let flags = self.trainer.read_nan_flags()?;
if flags[4] != 0 || flags[5] != 0 {
tracing::error!(
"NaN_DIAG step {} PRE_FORWARD: f32_params={} bf16_params={}",
self.steps_since_varmap_sync, flags[4], flags[5]
);
}
}
// ── Step 1: Spectral normalization BEFORE forward pass ─────────
self.trainer.apply_spectral_norm(&mut self.online_dueling, &mut self.online_branching)
.map_err(|e| anyhow::anyhow!("Spectral norm (pre-forward): {e}"))?;
if nan_diag {
self.trainer.run_nan_checks_pre_forward()?;
let flags = self.trainer.read_nan_flags()?;
if flags[4] != 0 || flags[5] != 0 {
tracing::error!(
"NaN_DIAG step {} POST_SPECTRAL: f32_params={} bf16_params={}",
self.steps_since_varmap_sync, flags[4], flags[5]
);
}
}
// ── Step 2: Upload batch + replay graph_forward ──────────────────
let _fused_placeholder = self.trainer.train_step_gpu(
gpu_batch,
@@ -666,6 +690,9 @@ impl FusedTrainingCtx {
&self.target_dueling, &self.target_branching,
).map_err(|e| anyhow::anyhow!("Fused train_step_gpu (forward only): {e}"))?;
// ── NaN detection: check key buffers after graph_forward ─────────
self.trainer.run_nan_checks_post_forward(self.batch_size)?;
// ── Step 2b: HER donor computation (outside graph_aux) ───────────
// Donor indices vary per step (random/future/final). The computation
// fills her.donor_indices GPU buffer. graph_aux captures the relabel
@@ -1419,6 +1446,11 @@ impl FusedTrainingCtx {
&self.trainer.grad_norm_buf
}
/// 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> {
self.trainer.read_nan_flags()
}
/// Get CUDA stream reference for DtoH transfers after compute_q_values.
pub(crate) fn stream(&self) -> &Arc<cudarc::driver::CudaStream> {
&self.stream

View File

@@ -1222,6 +1222,20 @@ impl DQNTrainer {
!guard_past_warmup,
).map_err(|e| anyhow::anyhow!("guard check: {e}"))?;
if gr.halt_nan {
// Read GPU-side NaN detection flags to identify source
if let Some(ref 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()
.filter(|(_, &f)| f != 0)
.map(|(i, _)| names[i])
.collect();
tracing::error!(
"NaN SOURCE at step {}: flagged=[{}] (0=mse_loss 1=grad_buf 2=d_val_logits 3=bf16_params)",
train_step_count, flagged.join(", ")
);
}
}
return Err(anyhow::anyhow!(
"NaN/Inf at step {}: loss={}, grad={}",
train_step_count, gr.raw_loss, gr.raw_grad_norm