diff --git a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu index 658927bc1..f735ab27f 100644 --- a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu +++ b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu @@ -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 diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 3254be1b2..ecb1d6e52 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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, /// #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::(b * bn_alloc_dim) .map_err(|e| MLError::ModelError(format!("alloc bn_d_hidden: {e}")))?; + let nan_flags_buf = stream.alloc_zeros::(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) -> 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, 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. diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index a5e569c23..ca3b47a32 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -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 { &self.stream diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 760ded631..1e0fb0d3f 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -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