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:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user